tidy codes
Kim committed
Mar 30, 2026 at 09:59 UTC
b811d09dc74ab227147f6ad7e8249907d541979b
3 files changed
+261
-282
portal/api_server.go
+47
-66
@@ -240,7 +240,11 @@ func (s *Server) handleRenew(w http.ResponseWriter, r *http.Request) {
240
return
241
}
242
243
- resp, err := s.renewLease(req, clientIP)
243
+ ttl := s.cfg.LeaseTTL
244
+ if req.TTL > 0 {
245
+ ttl = time.Duration(req.TTL) * time.Second
246
+ }
247
+ record, err := s.registry.Renew(strings.TrimSpace(req.LeaseID), req.ReverseToken, ttl, clientIP, utils.SanitizeReportedIP(req.ReportedIP))
248
if err != nil {
249
status, code := http.StatusBadRequest, types.APIErrorCodeInvalidRequest
250
if errors.Is(err, errLeaseNotFound) {
@@ -256,7 +260,7 @@ func (s *Server) handleRenew(w http.ResponseWriter, r *http.Request) {
260
return
261
}
262
259
- utils.WriteAPIData(w, http.StatusOK, resp)
263
+ utils.WriteAPIData(w, http.StatusOK, types.RenewResponse{LeaseID: record.ID, ExpiresAt: record.ExpiresAt})
264
}
265
266
func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
@@ -271,7 +275,8 @@ func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
275
return
276
}
277
274
- if err := s.unregisterLease(req); err != nil {
278
+ record, err := s.registry.Unregister(strings.TrimSpace(req.LeaseID), req.ReverseToken)
279
+ if err != nil {
280
status, code := http.StatusBadRequest, types.APIErrorCodeInvalidRequest
281
if errors.Is(err, errLeaseNotFound) {
282
status, code = http.StatusNotFound, types.APIErrorCodeLeaseNotFound
@@ -282,6 +287,9 @@ func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
287
utils.WriteAPIError(w, status, code, err.Error())
288
return
289
}
290
+ if record != nil {
291
+ record.Close()
292
+ }
293
294
utils.WriteAPIOK(w, http.StatusOK)
295
}
@@ -304,30 +312,22 @@ func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
312
return
313
}
314
307
- lease, err := s.registry.FindByID(leaseID)
308
- if err == nil && !s.registry.policy.IsLeaseRoutable(lease.ID) {
309
- err = errLeaseRejected
310
- }
311
- if err == nil && !utils.TokenMatches(lease.ReverseToken, token) {
312
- err = errUnauthorized
313
- }
314
- if err == nil && lease.stream == nil {
315
- err = errTransportMismatch
316
- }
317
- switch {
318
- case errors.Is(err, errLeaseNotFound):
315
+ lease, err := s.admitLeaseByID(leaseID, token, false)
316
+ switch err {
317
+ case nil:
318
+ case errLeaseNotFound:
319
utils.WriteAPIError(w, http.StatusNotFound, types.APIErrorCodeLeaseNotFound, err.Error())
320
return
321
- case errors.Is(err, errLeaseRejected):
321
+ case errLeaseRejected:
322
utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeLeaseRejected, "lease is not approved for routing")
323
return
324
- case errors.Is(err, errUnauthorized):
324
+ case errUnauthorized:
325
utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeUnauthorized, err.Error())
326
return
327
- case errors.Is(err, errTransportMismatch):
327
+ case errTransportMismatch:
328
utils.WriteAPIError(w, http.StatusConflict, types.APIErrorCodeTransportMismatch, "lease does not support stream transport")
329
return
330
- case err != nil:
330
+ default:
331
utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, err.Error())
332
return
333
}
@@ -396,34 +396,26 @@ func (s *Server) handleQUICTunnelConn(conn *quic.Conn) {
396
return
397
}
398
399
- lease, err := s.registry.FindByID(msg.LeaseID)
400
- if err == nil && !s.registry.policy.IsLeaseRoutable(lease.ID) {
401
- err = errLeaseRejected
402
- }
403
- if err == nil && !utils.TokenMatches(lease.ReverseToken, msg.ReverseToken) {
404
- err = errUnauthorized
405
- }
406
- if err == nil && (lease.stream == nil || lease.datagram == nil) {
407
- err = errTransportMismatch
408
- }
409
- switch {
410
- case errors.Is(err, errLeaseNotFound):
399
+ lease, err := s.admitLeaseByID(msg.LeaseID, msg.ReverseToken, true)
400
+ switch err {
401
+ case nil:
402
+ case errLeaseNotFound:
403
_ = json.NewEncoder(stream).Encode(types.QUICControlResponse{OK: false, Error: types.APIErrorCodeLeaseNotFound})
404
_ = conn.CloseWithError(1, "lease not found")
405
return
414
- case errors.Is(err, errUnauthorized):
415
- _ = json.NewEncoder(stream).Encode(types.QUICControlResponse{OK: false, Error: types.APIErrorCodeUnauthorized})
416
- _ = conn.CloseWithError(1, "unauthorized")
417
- return
418
- case errors.Is(err, errLeaseRejected):
406
+ case errLeaseRejected:
407
_ = json.NewEncoder(stream).Encode(types.QUICControlResponse{OK: false, Error: types.APIErrorCodeLeaseRejected})
408
_ = conn.CloseWithError(1, "lease rejected")
409
return
422
- case errors.Is(err, errTransportMismatch):
410
+ case errUnauthorized:
411
+ _ = json.NewEncoder(stream).Encode(types.QUICControlResponse{OK: false, Error: types.APIErrorCodeUnauthorized})
412
+ _ = conn.CloseWithError(1, "unauthorized")
413
+ return
414
+ case errTransportMismatch:
415
_ = json.NewEncoder(stream).Encode(types.QUICControlResponse{OK: false, Error: types.APIErrorCodeTransportMismatch})
416
_ = conn.CloseWithError(1, "transport mismatch")
417
return
426
- case err != nil:
418
+ default:
419
_ = json.NewEncoder(stream).Encode(types.QUICControlResponse{OK: false, Error: types.APIErrorCodeInvalidRequest})
420
_ = conn.CloseWithError(1, "invalid control message")
421
return
@@ -445,6 +437,23 @@ func (s *Server) handleQUICTunnelConn(conn *quic.Conn) {
437
Msg("quic tunnel connected")
438
}
439
440
+func (s *Server) admitLeaseByID(leaseID, token string, requireDatagram bool) (*leaseRecord, error) {
441
+ lease, err := s.registry.FindByID(leaseID)
442
+ if err != nil {
443
+ return nil, err
444
+ }
445
+ if !s.registry.policy.IsLeaseRoutable(lease.ID) {
446
+ return nil, errLeaseRejected
447
+ }
448
+ if !utils.TokenMatches(lease.ReverseToken, token) {
449
+ return nil, errUnauthorized
450
+ }
451
+ if lease.stream == nil || (requireDatagram && lease.datagram == nil) {
452
+ return nil, errTransportMismatch
453
+ }
454
+ return lease, nil
455
+}
456
+
457
func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (types.RegisterResponse, error) {
458
name, err := utils.NormalizeDNSLabel(req.Name)
459
if err != nil {
@@ -541,34 +550,6 @@ func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (type
550
return resp, nil
551
}
552
544
-func (s *Server) renewLease(req types.RenewRequest, clientIP string) (types.RenewResponse, error) {
545
- if s.registry.policy.IPFilter().IsIPBanned(clientIP) {
546
- return types.RenewResponse{}, errIPBanned
547
- }
548
-
549
- ttl := s.cfg.LeaseTTL
550
- if req.TTL > 0 {
551
- ttl = time.Duration(req.TTL) * time.Second
552
- }
553
- record, err := s.registry.Renew(strings.TrimSpace(req.LeaseID), req.ReverseToken, ttl, clientIP, utils.SanitizeReportedIP(req.ReportedIP))
554
- if err != nil {
555
- return types.RenewResponse{}, err
556
- }
557
-
558
- return types.RenewResponse{LeaseID: record.ID, ExpiresAt: record.ExpiresAt}, nil
559
-}
560
-
561
-func (s *Server) unregisterLease(req types.UnregisterRequest) error {
562
- record, err := s.registry.Unregister(strings.TrimSpace(req.LeaseID), req.ReverseToken)
563
- if err != nil {
564
- return err
565
- }
566
- if record != nil {
567
- record.Close()
568
- }
569
- return nil
570
-}
571
-
553
func (s *Server) runAPIServer() error {
554
err := s.apiServer.Serve(s.apiListener)
555
if err == nil || errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
portal/server.go
+25
-27
@@ -158,7 +158,6 @@ func NewServer(cfg ServerConfig) (*Server, error) {
158
Str("owner_private_key", ownerIdentity.PrivateKey).
159
Msg("generated relay owner private key; set OWNER_PRIVATE_KEY unique identity")
160
}
161
- cfg.OwnerPrivateKey = ""
161
162
runtime := policy.NewRuntime()
163
runtime.SetUDPPolicy(cfg.UDPPortCount > 0, 0)
@@ -491,32 +490,6 @@ func (s *Server) runSNIListener(ctx context.Context) error {
490
}
491
}
492
494
-func (s *Server) startOverlay() error {
495
- peerMux := http.NewServeMux()
496
- peerMux.HandleFunc(types.PathRoot, s.handleRoot)
497
- peerMux.HandleFunc(types.PathHealthz, s.handleHealthz)
498
- peerMux.HandleFunc(types.PathDiscovery, func(w http.ResponseWriter, r *http.Request) {
499
- if !s.DiscoveryEnabled() {
500
- http.NotFound(w, r)
501
- return
502
- }
503
- s.handleRelayDiscovery(w, r)
504
- })
505
-
506
- overlay, err := wireguard.NewOverlay(s.wgConfig, peerMux)
507
- if err != nil {
508
- return fmt.Errorf("start wireguard overlay: %w", err)
509
- }
510
-
511
- if err := overlay.Sync(s.cfg.PortalURL, s.relaySet.Snapshot()); err != nil {
512
- _ = overlay.Shutdown(context.Background())
513
- return fmt.Errorf("sync wireguard peers: %w", err)
514
- }
515
-
516
- s.overlay = overlay
517
- return nil
518
-}
519
-
493
func (s *Server) startQUICTunnelListener(apiTLS keyless.TLSMaterialConfig) error {
494
if len(apiTLS.KeyPEM) == 0 {
495
return fmt.Errorf("quic tunnel requires api tls key")
@@ -565,6 +538,31 @@ func (s *Server) runQUICTunnelListener(listener *quic.Listener) error {
538
}
539
}
540
541
+func (s *Server) startOverlay() error {
542
+ peerMux := http.NewServeMux()
543
+ peerMux.HandleFunc(types.PathRoot, s.handleRoot)
544
+ peerMux.HandleFunc(types.PathHealthz, s.handleHealthz)
545
+ peerMux.HandleFunc(types.PathDiscovery, func(w http.ResponseWriter, r *http.Request) {
546
+ if !s.DiscoveryEnabled() {
547
+ http.NotFound(w, r)
548
+ return
549
+ }
550
+ s.handleRelayDiscovery(w, r)
551
+ })
552
+
553
+ overlay, err := wireguard.NewOverlay(s.wgConfig, peerMux)
554
+ if err != nil {
555
+ return fmt.Errorf("start wireguard overlay: %w", err)
556
+ }
557
+
558
+ if err := overlay.Sync(s.cfg.PortalURL, s.relaySet.Snapshot()); err != nil {
559
+ _ = overlay.Shutdown(context.Background())
560
+ return fmt.Errorf("sync wireguard peers: %w", err)
561
+ }
562
+
563
+ s.overlay = overlay
564
+ return nil
565
+}
566
func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
567
ticker := time.NewTicker(defaultDiscoveryInterval)
568
defer ticker.Stop()
sdk/expose.go
+189
-189
@@ -133,6 +133,42 @@ func (e *Exposure) ActiveRelayURLs() []string {
133
return e.relaySet.ActiveRelayURLs()
134
}
135
136
+func (e *Exposure) Addr() net.Addr {
137
+ return listenerAddr("portal:exposure")
138
+}
139
+
140
+type exposureConn struct {
141
+ net.Conn
142
+ id uint64
143
+ localAddr string
144
+ remoteAddr string
145
+ closeOnce sync.Once
146
+}
147
+
148
+func (c *exposureConn) Close() error {
149
+ var closeErr error
150
+ c.closeOnce.Do(func() {
151
+ closeErr = c.Conn.Close()
152
+ if errors.Is(closeErr, net.ErrClosed) {
153
+ closeErr = nil
154
+ }
155
+
156
+ event := log.Info().
157
+ Uint64("conn_id", c.id).
158
+ Str("local_addr", c.localAddr).
159
+ Str("remote_addr", c.remoteAddr)
160
+ if closeErr != nil {
161
+ event = log.Warn().
162
+ Err(closeErr).
163
+ Uint64("conn_id", c.id).
164
+ Str("local_addr", c.localAddr).
165
+ Str("remote_addr", c.remoteAddr)
166
+ }
167
+ event.Msg("exposure connection closed")
168
+ })
169
+ return closeErr
170
+}
171
+
172
func (e *Exposure) Accept() (net.Conn, error) {
173
select {
174
case <-e.done:
@@ -158,129 +194,6 @@ func (e *Exposure) Accept() (net.Conn, error) {
194
}
195
}
196
161
-func (e *Exposure) Addr() net.Addr {
162
- return listenerAddr("portal:exposure")
163
-}
164
-
165
-func (e *Exposure) RunHTTP(ctx context.Context, handler http.Handler, localAddr string) error {
166
- var relayListener net.Listener
167
- e.listenerMu.RLock()
168
- activeListeners := make([]*Listener, 0, len(e.relayListeners))
169
- for _, relayURL := range e.relaySet.ActiveRelayURLs() {
170
- listener, ok := e.relayListeners[relayURL]
171
- if !ok {
172
- continue
173
- }
174
- activeListeners = append(activeListeners, listener)
175
- }
176
- e.listenerMu.RUnlock()
177
- if len(activeListeners) > 0 {
178
- relayListener = e
179
- }
180
- return RunHTTP(ctx, relayListener, handler, localAddr)
181
-}
182
-
183
-func RunHTTP(ctx context.Context, relayListener net.Listener, handler http.Handler, localAddr string) error {
184
- localAddr = strings.TrimSpace(localAddr)
185
-
186
- if relayListener == nil && localAddr == "" {
187
- return errors.New("relay listener or local address is required")
188
- }
189
-
190
- var relaySrv *http.Server
191
- if relayListener != nil {
192
- relaySrv = &http.Server{
193
- Handler: handler,
194
- ReadHeaderTimeout: defaultRequestTimeout,
195
- }
196
- }
197
-
198
- var localSrv *http.Server
199
- if localAddr != "" {
200
- localSrv = &http.Server{
201
- Addr: localAddr,
202
- Handler: handler,
203
- ReadHeaderTimeout: defaultRequestTimeout,
204
- }
205
- }
206
-
207
- serverCount := 0
208
- if relaySrv != nil {
209
- serverCount++
210
- }
211
- if localSrv != nil {
212
- serverCount++
213
- }
214
-
215
- results := make(chan error, serverCount)
216
- normalizeServeErr := func(err error, prefix string) error {
217
- if err == nil || errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
218
- return nil
219
- }
220
- return fmt.Errorf("%s: %w", prefix, err)
221
- }
222
-
223
- var (
224
- shutdownOnce sync.Once
225
- shutdownErr error
226
- )
227
- shutdown := func() error {
228
- shutdownOnce.Do(func() {
229
- shutdownCtx, cancel := context.WithTimeout(context.Background(), defaultHTTPShutdownTimeout)
230
- defer cancel()
231
-
232
- var localErr error
233
- if localSrv != nil {
234
- localErr = localSrv.Shutdown(shutdownCtx)
235
- if errors.Is(localErr, http.ErrServerClosed) {
236
- localErr = nil
237
- }
238
- }
239
-
240
- var relayErr error
241
- if relaySrv != nil {
242
- relayErr = relaySrv.Shutdown(shutdownCtx)
243
- if errors.Is(relayErr, http.ErrServerClosed) {
244
- relayErr = nil
245
- }
246
- }
247
-
248
- shutdownErr = errors.Join(localErr, relayErr)
249
- })
250
- return shutdownErr
251
- }
252
-
253
- if localSrv != nil {
254
- go func() {
255
- results <- normalizeServeErr(localSrv.ListenAndServe(), "serve local http")
256
- }()
257
- }
258
- if relaySrv != nil {
259
- go func() {
260
- results <- normalizeServeErr(relaySrv.Serve(relayListener), "serve relay http")
261
- }()
262
- }
263
-
264
- var serveErr error
265
- remaining := serverCount
266
- ctxDone := ctx.Done()
267
- for remaining > 0 {
268
- select {
269
- case err := <-results:
270
- remaining--
271
- if err != nil {
272
- serveErr = errors.Join(serveErr, err)
273
- _ = shutdown()
274
- }
275
- case <-ctxDone:
276
- _ = shutdown()
277
- ctxDone = nil
278
- }
279
- }
280
-
281
- return errors.Join(serveErr, shutdownErr)
282
-}
283
-
197
func (e *Exposure) Close() error {
198
var closeErr error
199
e.closeOnce.Do(func() {
@@ -317,6 +230,53 @@ func (e *Exposure) Close() error {
230
return closeErr
231
}
232
233
+const defaultDiscoveryInterval = 30 * time.Second
234
+
235
+func (e *Exposure) runRelayDiscoveryLoop(ctx context.Context) {
236
+ for {
237
+ relayURLs := append([]string(nil), e.relaySet.ActiveRelayURLs()...)
238
+ if len(relayURLs) > 0 {
239
+ var discoveredRelayURLs []string
240
+
241
+ for _, relayURL := range relayURLs {
242
+ resp, err := discovery.DiscoverRelayDiscovery(ctx, relayURL, e.rootCAPEM, nil)
243
+ if err != nil {
244
+ if ctx.Err() != nil {
245
+ return
246
+ }
247
+ continue
248
+ }
249
+
250
+ now := time.Now().UTC()
251
+ var descriptorRelayURLs []string
252
+ descriptorRelayURLs, _, _, _, err = e.relaySet.ApplyRelayDiscoveryResponse(relayURL, relayURL, resp, now)
253
+ if err != nil {
254
+ continue
255
+ }
256
+
257
+ if len(discoveredRelayURLs) == 0 {
258
+ discoveredRelayURLs = append([]string(nil), relayURLs...)
259
+ }
260
+ discoveredRelayURLs, err = utils.MergeRelayURLs(discoveredRelayURLs, nil, descriptorRelayURLs)
261
+ if err != nil {
262
+ continue
263
+ }
264
+ }
265
+
266
+ if len(discoveredRelayURLs) > 0 {
267
+ if e.relaySet == nil {
268
+ e.relaySet = discovery.NewRelaySet()
269
+ }
270
+ e.relaySet.ReplaceKnownRelayURLs(discoveredRelayURLs)
271
+ _ = e.reconcileRelayListeners(false)
272
+ }
273
+ }
274
+ if !utils.SleepOrDone(ctx, defaultDiscoveryInterval) {
275
+ return
276
+ }
277
+ }
278
+}
279
+
280
func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
281
if e.relaySet == nil {
282
e.relaySet = discovery.NewRelaySet()
@@ -454,38 +414,6 @@ func (e *Exposure) runListenerAcceptLoop(listener *Listener) {
414
}
415
}
416
457
-type exposureConn struct {
458
- net.Conn
459
- id uint64
460
- localAddr string
461
- remoteAddr string
462
- closeOnce sync.Once
463
-}
464
-
465
-func (c *exposureConn) Close() error {
466
- var closeErr error
467
- c.closeOnce.Do(func() {
468
- closeErr = c.Conn.Close()
469
- if errors.Is(closeErr, net.ErrClosed) {
470
- closeErr = nil
471
- }
472
-
473
- event := log.Info().
474
- Uint64("conn_id", c.id).
475
- Str("local_addr", c.localAddr).
476
- Str("remote_addr", c.remoteAddr)
477
- if closeErr != nil {
478
- event = log.Warn().
479
- Err(closeErr).
480
- Uint64("conn_id", c.id).
481
- Str("local_addr", c.localAddr).
482
- Str("remote_addr", c.remoteAddr)
483
- }
484
- event.Msg("exposure connection closed")
485
- })
486
- return closeErr
487
-}
488
-
417
func (e *Exposure) AcceptDatagram() (types.DatagramFrame, error) {
418
if !e.udpEnabled {
419
return types.DatagramFrame{}, net.ErrClosed
@@ -565,49 +493,121 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
493
}
494
}
495
568
-const defaultDiscoveryInterval = 30 * time.Second
496
+func (e *Exposure) RunHTTP(ctx context.Context, handler http.Handler, localAddr string) error {
497
+ var relayListener net.Listener
498
+ e.listenerMu.RLock()
499
+ activeListeners := make([]*Listener, 0, len(e.relayListeners))
500
+ for _, relayURL := range e.relaySet.ActiveRelayURLs() {
501
+ listener, ok := e.relayListeners[relayURL]
502
+ if !ok {
503
+ continue
504
+ }
505
+ activeListeners = append(activeListeners, listener)
506
+ }
507
+ e.listenerMu.RUnlock()
508
+ if len(activeListeners) > 0 {
509
+ relayListener = e
510
+ }
511
+ return RunHTTP(ctx, relayListener, handler, localAddr)
512
+}
513
570
-func (e *Exposure) runRelayDiscoveryLoop(ctx context.Context) {
571
- for {
572
- relayURLs := append([]string(nil), e.relaySet.ActiveRelayURLs()...)
573
- if len(relayURLs) > 0 {
574
- var discoveredRelayURLs []string
514
+func RunHTTP(ctx context.Context, relayListener net.Listener, handler http.Handler, localAddr string) error {
515
+ localAddr = strings.TrimSpace(localAddr)
516
576
- for _, relayURL := range relayURLs {
577
- resp, err := discovery.DiscoverRelayDiscovery(ctx, relayURL, e.rootCAPEM, nil)
578
- if err != nil {
579
- if ctx.Err() != nil {
580
- return
581
- }
582
- continue
583
- }
517
+ if relayListener == nil && localAddr == "" {
518
+ return errors.New("relay listener or local address is required")
519
+ }
520
585
- now := time.Now().UTC()
586
- var descriptorRelayURLs []string
587
- descriptorRelayURLs, _, _, _, err = e.relaySet.ApplyRelayDiscoveryResponse(relayURL, relayURL, resp, now)
588
- if err != nil {
589
- continue
590
- }
521
+ var relaySrv *http.Server
522
+ if relayListener != nil {
523
+ relaySrv = &http.Server{
524
+ Handler: handler,
525
+ ReadHeaderTimeout: defaultRequestTimeout,
526
+ }
527
+ }
528
592
- if len(discoveredRelayURLs) == 0 {
593
- discoveredRelayURLs = append([]string(nil), relayURLs...)
594
- }
595
- discoveredRelayURLs, err = utils.MergeRelayURLs(discoveredRelayURLs, nil, descriptorRelayURLs)
596
- if err != nil {
597
- continue
529
+ var localSrv *http.Server
530
+ if localAddr != "" {
531
+ localSrv = &http.Server{
532
+ Addr: localAddr,
533
+ Handler: handler,
534
+ ReadHeaderTimeout: defaultRequestTimeout,
535
+ }
536
+ }
537
+
538
+ serverCount := 0
539
+ if relaySrv != nil {
540
+ serverCount++
541
+ }
542
+ if localSrv != nil {
543
+ serverCount++
544
+ }
545
+
546
+ results := make(chan error, serverCount)
547
+ normalizeServeErr := func(err error, prefix string) error {
548
+ if err == nil || errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
549
+ return nil
550
+ }
551
+ return fmt.Errorf("%s: %w", prefix, err)
552
+ }
553
+
554
+ var (
555
+ shutdownOnce sync.Once
556
+ shutdownErr error
557
+ )
558
+ shutdown := func() error {
559
+ shutdownOnce.Do(func() {
560
+ shutdownCtx, cancel := context.WithTimeout(context.Background(), defaultHTTPShutdownTimeout)
561
+ defer cancel()
562
+
563
+ var localErr error
564
+ if localSrv != nil {
565
+ localErr = localSrv.Shutdown(shutdownCtx)
566
+ if errors.Is(localErr, http.ErrServerClosed) {
567
+ localErr = nil
568
}
569
}
570
601
- if len(discoveredRelayURLs) > 0 {
602
- if e.relaySet == nil {
603
- e.relaySet = discovery.NewRelaySet()
571
+ var relayErr error
572
+ if relaySrv != nil {
573
+ relayErr = relaySrv.Shutdown(shutdownCtx)
574
+ if errors.Is(relayErr, http.ErrServerClosed) {
575
+ relayErr = nil
576
}
605
- e.relaySet.ReplaceKnownRelayURLs(discoveredRelayURLs)
606
- _ = e.reconcileRelayListeners(false)
577
}
608
- }
609
- if !utils.SleepOrDone(ctx, defaultDiscoveryInterval) {
610
- return
578
+
579
+ shutdownErr = errors.Join(localErr, relayErr)
580
+ })
581
+ return shutdownErr
582
+ }
583
+
584
+ if localSrv != nil {
585
+ go func() {
586
+ results <- normalizeServeErr(localSrv.ListenAndServe(), "serve local http")
587
+ }()
588
+ }
589
+ if relaySrv != nil {
590
+ go func() {
591
+ results <- normalizeServeErr(relaySrv.Serve(relayListener), "serve relay http")
592
+ }()
593
+ }
594
+
595
+ var serveErr error
596
+ remaining := serverCount
597
+ ctxDone := ctx.Done()
598
+ for remaining > 0 {
599
+ select {
600
+ case err := <-results:
601
+ remaining--
602
+ if err != nil {
603
+ serveErr = errors.Join(serveErr, err)
604
+ _ = shutdown()
605
+ }
606
+ case <-ctxDone:
607
+ _ = shutdown()
608
+ ctxDone = nil
609
}
610
}
611
+
612
+ return errors.Join(serveErr, shutdownErr)
613
}