chore: remove ctx in struct
Kim committed
Mar 11, 2026 at 10:18 UTC
b906cbc5a22044f3cba0c00a64b254eaaa680cf0
3 files changed
+63
-107
portal/server.go
+14
-39
@@ -52,7 +52,6 @@ type Server struct {
52
apiTLSClose io.Closer
53
apiListener net.Listener
54
apiServer *http.Server
55
- ctx context.Context
55
cancel context.CancelFunc
56
group *errgroup.Group
57
routes *routeTable
@@ -160,14 +159,13 @@ func (s *Server) Start(ctx context.Context) error {
159
s.sniListener = sniListener
160
s.apiServer = apiServer
161
s.apiTLSClose = apiCloser
163
- s.ctx = groupCtx
162
s.cancel = cancel
163
s.group = group
164
165
group.Go(s.runAPIServer)
168
- group.Go(s.runSNIListener)
169
- group.Go(s.runLeaseJanitor)
170
- group.Go(s.watchContext)
166
+ group.Go(func() error { return s.runSNIListener(groupCtx) })
167
+ group.Go(func() error { return s.runLeaseJanitor(groupCtx) })
168
+ group.Go(func() error { return s.watchContext(groupCtx) })
169
return nil
170
}
171
@@ -243,20 +241,20 @@ func (s *Server) ListLeases() []LeaseSnapshot {
241
return out
242
}
243
246
-func (s *Server) runSNIListener() error {
244
+func (s *Server) runSNIListener(ctx context.Context) error {
245
for {
246
conn, err := s.sniListener.Accept()
247
if err != nil {
250
- if errors.Is(err, net.ErrClosed) || s.isClosed() {
248
+ if errors.Is(err, net.ErrClosed) || ctx.Err() != nil {
249
return nil
250
}
251
return err
252
}
255
- go s.handleSNIConn(conn)
253
+ go s.handleSNIConn(ctx, conn)
254
}
255
}
256
259
-func (s *Server) handleSNIConn(conn net.Conn) {
257
+func (s *Server) handleSNIConn(ctx context.Context, conn net.Conn) {
258
clientHello, wrappedConn, err := l4.InspectClientHello(conn, s.cfg.ClientHelloTimeout)
259
if err != nil {
260
_ = wrappedConn.Close()
@@ -270,7 +268,7 @@ func (s *Server) handleSNIConn(conn net.Conn) {
268
}
269
270
if serverName == s.cfg.RootHost && s.cfg.RootFallbackAddr != "" {
273
- s.bridgeToFallback(wrappedConn)
271
+ s.bridgeToFallback(ctx, wrappedConn)
272
return
273
}
274
@@ -292,7 +290,7 @@ func (s *Server) handleSNIConn(conn net.Conn) {
290
return
291
}
292
295
- claimCtx, cancel := context.WithTimeout(s.context(), s.cfg.ClaimTimeout)
293
+ claimCtx, cancel := context.WithTimeout(ctx, s.cfg.ClaimTimeout)
294
defer cancel()
295
296
session, err := record.Broker.Claim(claimCtx)
@@ -305,9 +303,9 @@ func (s *Server) handleSNIConn(conn net.Conn) {
303
_ = session.Close()
304
}
305
308
-func (s *Server) bridgeToFallback(conn net.Conn) {
306
+func (s *Server) bridgeToFallback(ctx context.Context, conn net.Conn) {
307
dialer := &net.Dialer{Timeout: 5 * time.Second}
310
- upstream, err := dialer.DialContext(s.context(), "tcp", HostPortOrLoopback(s.cfg.RootFallbackAddr))
308
+ upstream, err := dialer.DialContext(ctx, "tcp", HostPortOrLoopback(s.cfg.RootFallbackAddr))
309
if err != nil {
310
_ = conn.Close()
311
return
@@ -315,11 +313,10 @@ func (s *Server) bridgeToFallback(conn net.Conn) {
313
bridgeConns(conn, upstream)
314
}
315
318
-func (s *Server) runLeaseJanitor() error {
316
+func (s *Server) runLeaseJanitor(ctx context.Context) error {
317
ticker := time.NewTicker(5 * time.Second)
318
defer ticker.Stop()
319
322
- ctx := s.context()
320
for {
321
select {
322
case <-ctx.Done():
@@ -350,11 +347,8 @@ func (s *Server) cleanupExpiredLeases() {
347
}
348
}
349
353
-func (s *Server) watchContext() error {
354
- if s.ctx == nil {
355
- return nil
356
- }
357
- <-s.ctx.Done()
350
+func (s *Server) watchContext(ctx context.Context) error {
351
+ <-ctx.Done()
352
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
353
defer cancel()
354
return s.Shutdown(shutdownCtx)
@@ -387,25 +381,6 @@ func closeWrite(conn net.Conn) {
381
}
382
}
383
390
-func (s *Server) context() context.Context {
391
- if s.ctx != nil {
392
- return s.ctx
393
- }
394
- return context.Background()
395
-}
396
-
397
-func (s *Server) isClosed() bool {
398
- if s.ctx == nil {
399
- return false
400
- }
401
- select {
402
- case <-s.ctx.Done():
403
- return true
404
- default:
405
- return false
406
- }
407
-}
408
-
384
func (s *Server) snapshotForLease(record *leaseRecord) LeaseSnapshot {
385
if record == nil {
386
return LeaseSnapshot{}
sdk/listener.go
+48
-67
@@ -41,7 +41,7 @@ type Listener struct {
41
leaseTTL time.Duration
42
renewBefore time.Duration
43
handshakeTimeout time.Duration
44
- ctx context.Context
44
+ doneCh <-chan struct{}
45
cancel context.CancelFunc
46
api *apiClient
47
accepted chan net.Conn
@@ -88,7 +88,7 @@ func NewListener(ctx context.Context, relayURL string, cfg ListenerConfig) (*Lis
88
}
89
90
l := &Listener{
91
- ctx: listenerCtx,
91
+ doneCh: listenerCtx.Done(),
92
cancel: cancel,
93
api: api,
94
accepted: make(chan net.Conn, max(readyTarget*2, 1)),
@@ -132,15 +132,15 @@ func NewListener(ctx context.Context, relayURL string, cfg ListenerConfig) (*Lis
132
l.mu.Unlock()
133
134
for i := 0; i < l.readyTarget; i++ {
135
- go l.runSessionLoop()
135
+ go l.runSessionLoop(listenerCtx)
136
}
137
- go l.runRenewLoop()
137
+ go l.runRenewLoop(listenerCtx)
138
return l, nil
139
}
140
141
func (l *Listener) Accept() (net.Conn, error) {
142
select {
143
- case <-l.ctx.Done():
143
+ case <-l.doneCh:
144
return nil, net.ErrClosed
145
case conn := <-l.accepted:
146
return conn, nil
@@ -219,11 +219,11 @@ func (l *Listener) PublicURLs() []string {
219
return urls
220
}
221
222
-func (l *Listener) runSessionLoop() {
222
+func (l *Listener) runSessionLoop(ctx context.Context) {
223
var retries int
224
225
for {
226
- claimed, err := l.runSession()
226
+ claimed, err := l.runSession(ctx)
227
switch {
228
case err == nil:
229
retries = 0
@@ -235,14 +235,14 @@ func (l *Listener) runSessionLoop() {
235
retries = 0
236
default:
237
retries++
238
- if !l.retryOrClose("reverse session connect", err, retries) {
238
+ if !l.retryOrClose(ctx, "reverse session connect", err, retries) {
239
return
240
}
241
}
242
}
243
}
244
245
-func (l *Listener) runRenewLoop() {
245
+func (l *Listener) runRenewLoop(ctx context.Context) {
246
interval := l.leaseTTL / 2
247
if interval <= 0 {
248
interval = 30 * time.Second
@@ -255,55 +255,43 @@ func (l *Listener) runRenewLoop() {
255
}
256
257
for {
258
- sleepOrDone(l.context(), interval)
259
- if l.isClosed() {
258
+ if !sleepOrDone(ctx, interval) {
259
return
260
}
261
262
var retries int
263
for {
265
- err := l.renewLease()
266
- switch {
267
- case err == nil:
268
- goto nextRenew
269
- case errors.Is(err, context.Canceled), errors.Is(err, net.ErrClosed):
264
+ err := l.renewLease(ctx)
265
+ if err == nil {
266
+ break
267
+ }
268
+ if errors.Is(err, context.Canceled) || errors.Is(err, net.ErrClosed) {
269
return
271
- default:
272
- retries++
273
- if !l.retryOrClose("lease renewal", err, retries) {
274
- return
275
- }
270
}
277
- }
271
279
- nextRenew:
272
+ retries++
273
+ if !l.retryOrClose(ctx, "lease renewal", err, retries) {
274
+ return
275
+ }
276
+ }
277
}
278
}
279
283
-func (l *Listener) runSession() (bool, error) {
284
- sessionCtx := l.context()
280
+func (l *Listener) runSession(ctx context.Context) (bool, error) {
281
l.mu.Lock()
282
leaseID := l.leaseID
283
l.mu.Unlock()
284
289
- conn, err := l.api.openReverseSession(sessionCtx, leaseID)
285
+ conn, err := l.api.openReverseSession(ctx, leaseID)
286
if err != nil {
287
return false, err
288
}
289
294
- claimed, err := l.awaitActivation(conn)
295
- if err != nil {
296
- _ = conn.Close()
297
- return claimed, err
298
- }
299
- return claimed, nil
300
-}
301
-
302
-func (l *Listener) awaitActivation(conn net.Conn) (bool, error) {
290
var marker [1]byte
291
for {
292
_ = conn.SetReadDeadline(time.Now().Add(2 * l.handshakeTimeout))
293
if _, err := io.ReadFull(conn, marker[:]); err != nil {
294
+ _ = conn.Close()
295
return false, err
296
}
297
_ = conn.SetReadDeadline(time.Time{})
@@ -312,44 +300,46 @@ func (l *Listener) awaitActivation(conn net.Conn) (bool, error) {
300
case types.MarkerKeepalive:
301
continue
302
case types.MarkerTLSStart:
315
- if err := l.activate(conn); err != nil {
303
+ if err := l.activate(ctx, conn); err != nil {
304
+ _ = conn.Close()
305
return true, err
306
}
307
return true, nil
308
default:
309
+ _ = conn.Close()
310
return false, fmt.Errorf("unexpected reverse marker: 0x%02x", marker[0])
311
}
312
}
313
}
314
325
-func (l *Listener) activate(conn net.Conn) error {
315
+func (l *Listener) activate(ctx context.Context, conn net.Conn) error {
316
l.mu.Lock()
317
tlsCfg := l.tlsConfig
318
l.mu.Unlock()
319
320
tlsConn := tls.Server(conn, tlsCfg)
331
- handshakeCtx, cancel := context.WithTimeout(l.context(), l.handshakeTimeout)
321
+ handshakeCtx, cancel := context.WithTimeout(ctx, l.handshakeTimeout)
322
defer cancel()
323
if err := tlsConn.HandshakeContext(handshakeCtx); err != nil {
324
return err
325
}
326
327
select {
338
- case <-l.ctx.Done():
328
+ case <-ctx.Done():
329
_ = tlsConn.Close()
340
- return l.context().Err()
330
+ return ctx.Err()
331
case l.accepted <- tlsConn:
332
return nil
333
}
334
}
335
346
-func (l *Listener) renewLease() error {
336
+func (l *Listener) renewLease(ctx context.Context) error {
337
l.mu.Lock()
338
leaseID := l.leaseID
339
l.mu.Unlock()
340
351
- ctx, cancel := context.WithTimeout(l.context(), 10*time.Second)
352
- err := l.api.renewLease(ctx, leaseID, l.leaseTTL)
341
+ requestCtx, cancel := context.WithTimeout(ctx, 10*time.Second)
342
+ err := l.api.renewLease(requestCtx, leaseID, l.leaseTTL)
343
cancel()
344
if err == nil {
345
return nil
@@ -364,7 +354,7 @@ func (l *Listener) renewLease() error {
354
Str("lease_id", leaseID).
355
Msg("lease not found on relay, attempting re-registration")
356
367
- if err := l.reregister(); err != nil {
357
+ if err := l.reregister(ctx); err != nil {
358
return err
359
}
360
@@ -376,29 +366,29 @@ func (l *Listener) renewLease() error {
366
return nil
367
}
368
379
-func (l *Listener) reregister() error {
380
- ctx, cancel := context.WithTimeout(l.context(), 10*time.Second)
369
+func (l *Listener) reregister(ctx context.Context) error {
370
+ requestCtx, cancel := context.WithTimeout(ctx, 10*time.Second)
371
defer cancel()
372
373
l.mu.Lock()
374
hostnames := append([]string(nil), l.hostnames...)
375
l.mu.Unlock()
376
387
- resp, err := l.api.registerLease(ctx, hostnames, l.leaseTTL)
377
+ resp, err := l.api.registerLease(requestCtx, hostnames, l.leaseTTL)
378
if err != nil {
379
return err
380
}
381
382
tlsConf, tlsCloser, err := keyless.BuildClientTLSConfig(l.api.baseURL.String(), resp.Hostnames)
383
if err != nil {
394
- _ = l.api.unregisterLease(ctx, resp.LeaseID)
384
+ _ = l.api.unregisterLease(requestCtx, resp.LeaseID)
385
return err
386
}
387
398
- if l.isClosed() {
388
+ if ctx.Err() != nil {
389
_ = l.api.unregisterLease(context.Background(), resp.LeaseID)
390
_ = tlsCloser.Close()
401
- return context.Canceled
391
+ return ctx.Err()
392
}
393
394
l.mu.Lock()
@@ -420,8 +410,8 @@ func isLeaseNotFound(err error) bool {
410
return errors.Is(err, &types.APIRequestError{Code: types.APIErrorCodeLeaseNotFound})
411
}
412
423
-func (l *Listener) retryOrClose(operation string, err error, retries int) bool {
424
- if l.isClosed() {
413
+func (l *Listener) retryOrClose(ctx context.Context, operation string, err error, retries int) bool {
414
+ if ctx.Err() != nil {
415
return false
416
}
417
@@ -447,16 +437,17 @@ func (l *Listener) retryOrClose(operation string, err error, retries int) bool {
437
Dur("retry_wait", l.retryWait).
438
Msg("operation failed; retrying")
439
450
- sleepOrDone(l.context(), l.retryWait)
451
- return !l.isClosed()
440
+ return sleepOrDone(ctx, l.retryWait)
441
}
442
454
-func sleepOrDone(ctx context.Context, d time.Duration) {
443
+func sleepOrDone(ctx context.Context, d time.Duration) bool {
444
timer := time.NewTimer(d)
445
defer timer.Stop()
446
select {
447
case <-ctx.Done():
448
+ return false
449
case <-timer.C:
450
+ return true
451
}
452
}
453
@@ -465,19 +456,9 @@ type listenerAddr string
456
func (a listenerAddr) Network() string { return "portal" }
457
func (a listenerAddr) String() string { return string(a) }
458
468
-func (l *Listener) context() context.Context {
469
- if l.ctx != nil {
470
- return l.ctx
471
- }
472
- return context.Background()
473
-}
474
-
475
-func (l *Listener) isClosed() bool {
476
- if l.ctx == nil {
477
- return false
478
- }
459
+func (l *Listener) done() bool {
460
select {
480
- case <-l.ctx.Done():
461
+ case <-l.doneCh:
462
return true
463
default:
464
return false
sdk/sdk_test.go
+1
-1
@@ -238,7 +238,7 @@ func TestNewListenerClosesAfterReverseSessionRetryBudgetExhausted(t *testing.T)
238
defer listener.Close()
239
240
waitForSDKTest(t, func() bool {
241
- return listener.isClosed()
241
+ return listener.done()
242
})
243
if connectCount.Load() < 2 {
244
t.Fatalf("connect count = %d, want at least 2", connectCount.Load())