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())