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 }