refactor: ensure safe shutdown using sync.Once for stop channels

- Renamed stopCh to stopClientCh in RelayClient struct for clarity. - Added sync.Once fields to RelayClient and rdRelay to prevent multiple closes of stop channels, avoiding potential panics. - Improved Close() methods in RelayClient and RDClient by collecting resources under lock then processing outside to enhance thread safety. - Added deferred close on incomingConnCh in leaseListenWorker for proper cleanup. - Reduced parameters in healthCheckWorker to use global config access.

lemon-mint committed Oct 30, 2025 at 14:44 UTC 42618ac18837c181bed6536af097ef793205f3b3
2 files changed +145 -178
portal/client.go
+12 -8
@@ -44,8 +44,9 @@ type RelayClient struct {
44 leases map[string]*leaseWithCred
45 leasesMu sync.Mutex
46
47 - stopCh chan struct{}
48 - waitGroup sync.WaitGroup
47 + stopClientCh chan struct{}
48 + stopOnce sync.Once // Ensure stopClientCh is closed only once
49 + waitGroup sync.WaitGroup
50
51 incommingConnCh chan *IncommingConn
52 }
@@ -76,7 +77,7 @@ func NewRelayClient(conn io.ReadWriteCloser) *RelayClient {
77 conn: conn,
78 sess: sess,
79 leases: make(map[string]*leaseWithCred),
79 - stopCh: make(chan struct{}),
80 + stopClientCh: make(chan struct{}),
81 incommingConnCh: make(chan *IncommingConn),
82 }
83
@@ -96,8 +97,10 @@ func (g *RelayClient) Ping() (time.Duration, error) {
97 func (g *RelayClient) Close() error {
98 log.Debug().Msg("[RelayClient] Closing relay client")
99
99 - // Signal workers to stop
100 - close(g.stopCh)
100 + // Signal workers to stop (only once)
101 + g.stopOnce.Do(func() {
102 + close(g.stopClientCh)
103 + })
104
105 var errs []error
106
@@ -137,7 +140,7 @@ func (g *RelayClient) leaseUpdateWorker() {
140 defer ticker.Stop()
141 for {
142 select {
140 - case <-g.stopCh:
143 + case <-g.stopClientCh:
144 return
145 case <-ticker.C:
146 clear(updateRequired)
@@ -163,11 +166,12 @@ func (g *RelayClient) leaseUpdateWorker() {
166
167 func (g *RelayClient) leaseListenWorker() {
168 defer g.waitGroup.Done()
169 + defer close(g.incommingConnCh)
170 log.Debug().Msg("[RelayClient] Lease listen worker started")
171
172 for {
173 select {
170 - case <-g.stopCh:
174 + case <-g.stopClientCh:
175 log.Debug().Msg("[RelayClient] Lease listen worker stopped")
176 return
177 default:
@@ -181,7 +185,7 @@ func (g *RelayClient) leaseListenWorker() {
185 if err != nil {
186 // Check if we're supposed to stop
187 select {
184 - case <-g.stopCh:
188 + case <-g.stopClientCh:
189 return
190 default:
191 log.Debug().Err(err).Msg("[RelayClient] Error accepting stream, retrying")
sdk/sdk.go
+133 -170
@@ -14,6 +14,7 @@ import (
14 "github.com/gorilla/websocket"
15 "github.com/gosuda/portal/portal"
16 "github.com/gosuda/portal/portal/core/cryptoops"
17 + "github.com/gosuda/portal/portal/core/proto/rdsec"
18 "github.com/gosuda/portal/portal/core/proto/rdverb"
19 "github.com/gosuda/portal/portal/utils/wsstream"
20 "github.com/rs/zerolog/log"
@@ -94,11 +95,12 @@ func WithReconnectInterval(interval time.Duration) Option {
95 }
96
97 type rdRelay struct {
97 - addr string
98 - client *portal.RelayClient
99 - dialer func(context.Context, string) (io.ReadWriteCloser, error)
100 - stop chan struct{}
101 - mu sync.Mutex
98 + addr string
99 + client *portal.RelayClient
100 + dialer func(context.Context, string) (io.ReadWriteCloser, error)
101 + stop chan struct{}
102 + stopOnce sync.Once // Ensure stop channel is closed only once
103 + mu sync.Mutex
104 }
105
106 var _ net.Conn = (*RDConnection)(nil)
@@ -174,6 +176,7 @@ type RDClient struct {
176 config *RDClientConfig
177
178 stopch chan struct{}
179 + stopOnce sync.Once // Ensure stopch is closed only once
180 waitGroup sync.WaitGroup // Track all listener workers
181 }
182
@@ -211,34 +214,12 @@ func NewClient(opt ...Option) (*RDClient, error) {
214 // Initialize relays from bootstrap servers
215 var connectionErrors []error
216 for _, server := range config.BootstrapServers {
214 - log.Debug().Str("server", server).Msg("[SDK] Connecting to bootstrap server")
215 - conn, err := config.Dialer(context.Background(), server)
217 + err := client.AddRelay(server, config.Dialer)
218 if err != nil {
219 log.Error().Err(err).Str("server", server).Msg("[SDK] Failed to connect to bootstrap server")
220 connectionErrors = append(connectionErrors, err)
219 - continue // Skip failed connections
221 }
221 -
222 - relayClient := portal.NewRelayClient(conn)
223 - if relayClient == nil {
224 - log.Error().Str("server", server).Msg("[SDK] Failed to create relay client")
225 - conn.Close()
226 - connectionErrors = append(connectionErrors, ErrFailedToCreateClient)
227 - continue
228 - }
229 -
222 log.Debug().Str("server", server).Msg("[SDK] Successfully connected to bootstrap server")
231 - relay := &rdRelay{
232 - addr: server,
233 - client: relayClient,
234 - dialer: config.Dialer,
235 - stop: make(chan struct{}),
236 - }
237 - client.relays[server] = relay
238 -
239 - // Start health monitoring for this relay
240 - client.waitGroup.Add(1)
241 - go client.healthCheckWorker(relay, config)
223 }
224
225 // If no relays were successfully connected, return an error
@@ -356,6 +337,10 @@ func (g *RDClient) Listen(cred *cryptoops.Credential, name string, alpns []strin
337 }
338
339 lease := &rdverb.Lease{
340 + Identity: &rdsec.Identity{
341 + Id: cred.ID(),
342 + PublicKey: cred.PublicKey(),
343 + },
344 Name: name,
345 Alpn: alpns,
346 }
@@ -472,32 +457,46 @@ func (g *RDClient) listenerWorker(server *rdRelay) {
457 }
458
459 func (g *RDClient) Close() error {
460 + log.Debug().Msg("[SDK] Closing RDClient")
461 var errs []error
462
477 - // Signal all goroutines to stop
478 - close(g.stopch)
463 + // Signal all goroutines to stop (only once)
464 + g.stopOnce.Do(func() {
465 + close(g.stopch)
466 + })
467
468 g.mu.Lock()
469 + listeners := make([]*RDListener, 0, len(g.listeners))
470 for _, listener := range g.listeners {
471 + listeners = append(listeners, listener)
472 + }
473 + relays := make([]*rdRelay, 0, len(g.relays))
474 + for _, relay := range g.relays {
475 + relays = append(relays, relay)
476 + }
477 + g.mu.Unlock()
478 +
479 + // Close all listeners first
480 + for _, listener := range listeners {
481 if err := listener.Close(); err != nil {
482 + log.Error().Err(err).Msg("[SDK] Error closing listener")
483 errs = append(errs, err)
484 }
485 }
486 - g.listeners = make(map[string]*RDListener)
486
487 // Stop all relays
489 - for _, server := range g.relays {
490 - close(server.stop) // Signal relay goroutines to stop
491 - if err := server.client.Close(); err != nil {
488 + for _, relay := range relays {
489 + if err := g.RemoveRelay(relay.addr); err != nil && err != ErrRelayNotFound {
490 + log.Error().Err(err).Str("relay", relay.addr).Msg("[SDK] Error removing relay")
491 errs = append(errs, err)
492 }
493 }
495 - g.relays = make(map[string]*rdRelay)
496 - g.mu.Unlock()
494
495 // Wait for all listener workers to finish
496 + log.Debug().Msg("[SDK] Waiting for all workers to finish")
497 g.waitGroup.Wait()
498
499 + log.Debug().Msg("[SDK] RDClient closed successfully")
500 if len(errs) > 0 {
501 return errs[0]
502 }
@@ -505,10 +504,10 @@ func (g *RDClient) Close() error {
504 }
505
506 // healthCheckWorker periodically checks relay health and reconnects if needed
508 -func (g *RDClient) healthCheckWorker(relay *rdRelay, config *RDClientConfig) {
507 +func (g *RDClient) healthCheckWorker(relay *rdRelay) {
508 defer g.waitGroup.Done()
509
511 - ticker := time.NewTicker(config.HealthCheckInterval)
510 + ticker := time.NewTicker(g.config.HealthCheckInterval)
511 defer ticker.Stop()
512
513 log.Debug().Str("relay", relay.addr).Msg("[SDK] Health check worker started")
@@ -516,19 +515,28 @@ func (g *RDClient) healthCheckWorker(relay *rdRelay, config *RDClientConfig) {
515 for {
516 select {
517 case <-g.stopch:
519 - log.Debug().Str("relay", relay.addr).Msg("[SDK] Health check worker stopped")
518 + log.Debug().Str("relay", relay.addr).Msg("[SDK] Health check worker stopped (client closing)")
519 return
520 case <-relay.stop:
522 - log.Debug().Str("relay", relay.addr).Msg("[SDK] Relay stopped, health check worker exiting")
521 + log.Debug().Str("relay", relay.addr).Msg("[SDK] Health check worker stopped (relay stopped)")
522 return
523 case <-ticker.C:
524 + // Check if client is still active
525 + select {
526 + case <-g.stopch:
527 + return
528 + case <-relay.stop:
529 + return
530 + default:
531 + }
532 +
533 relay.mu.Lock()
534 client := relay.client
535 relay.mu.Unlock()
536
537 if client == nil {
538 log.Warn().Str("relay", relay.addr).Msg("[SDK] Relay client is nil, attempting reconnection")
531 - g.reconnectRelay(relay, config)
539 + g.reconnectRelay(relay)
540 continue
541 }
542
@@ -539,132 +547,83 @@ func (g *RDClient) healthCheckWorker(relay *rdRelay, config *RDClientConfig) {
547 Err(err).
548 Str("relay", relay.addr).
549 Msg("[SDK] Health check failed, attempting reconnection")
542 - g.reconnectRelay(relay, config)
543 - } else {
544 - log.Debug().Str("relay", relay.addr).Msg("[SDK] Health check passed")
550 + g.reconnectRelay(relay)
551 + // Continue monitoring instead of returning
552 + continue
553 }
554 }
555 }
556 }
557
558 // reconnectRelay attempts to reconnect to a relay server
551 -func (g *RDClient) reconnectRelay(relay *rdRelay, config *RDClientConfig) {
552 - relay.mu.Lock()
559 +func (g *RDClient) reconnectRelay(relay *rdRelay) {
560 + addr := relay.addr
561 + dialer := relay.dialer
562
554 - // Close old client if exists
555 - if relay.client != nil {
556 - log.Debug().Str("relay", relay.addr).Msg("[SDK] Closing old relay client")
557 - relay.client.Close()
558 - relay.client = nil
559 - }
560 - relay.mu.Unlock()
563 + log.Debug().Str("relay", addr).Msg("[SDK] Starting reconnection process")
564
562 - maxRetries := config.ReconnectMaxRetries
563 - if maxRetries == 0 {
564 - maxRetries = -1 // Infinite retries
565 + // Remove the failed relay
566 + if err := g.RemoveRelay(addr); err != nil && err != ErrRelayNotFound {
567 + log.Error().Err(err).Str("relay", addr).Msg("[SDK] Error removing relay during reconnection")
568 }
569
567 - attempt := 0
568 - for {
569 - // Check if we should stop
570 - select {
571 - case <-g.stopch:
572 - log.Debug().Str("relay", relay.addr).Msg("[SDK] Client stopped, abandoning reconnection")
573 - return
574 - case <-relay.stop:
575 - log.Debug().Str("relay", relay.addr).Msg("[SDK] Relay stopped, abandoning reconnection")
576 - return
577 - default:
578 - }
579 -
580 - attempt++
581 - if maxRetries > 0 && attempt > maxRetries {
582 - log.Error().
583 - Str("relay", relay.addr).
584 - Int("attempts", attempt-1).
585 - Msg("[SDK] Max reconnection attempts reached, giving up")
586 - return
587 - }
588 -
589 - log.Debug().
590 - Str("relay", relay.addr).
591 - Int("attempt", attempt).
592 - Msg("[SDK] Attempting to reconnect")
570 + // Start reconnection in a goroutine
571 + g.waitGroup.Add(1)
572 + go func() {
573 + defer g.waitGroup.Done()
574
594 - // Attempt to connect
595 - conn, err := relay.dialer(context.Background(), relay.addr)
596 - if err != nil {
597 - log.Warn().
598 - Err(err).
599 - Str("relay", relay.addr).
600 - Int("attempt", attempt).
601 - Msg("[SDK] Reconnection attempt failed")
575 + retries := 0
576 + maxRetries := g.config.ReconnectMaxRetries
577
603 - // Wait before next retry
578 + for {
579 + // Check if client is shutting down
580 select {
581 case <-g.stopch:
582 + log.Debug().Str("relay", addr).Msg("[SDK] Reconnection cancelled (client closing)")
583 return
607 - case <-relay.stop:
608 - return
609 - case <-time.After(config.ReconnectInterval):
610 - continue
584 + default:
585 }
612 - }
586
614 - // Create new relay client
615 - relayClient := portal.NewRelayClient(conn)
616 - if relayClient == nil {
617 - log.Error().Str("relay", relay.addr).Msg("[SDK] Failed to create relay client after reconnection")
618 - conn.Close()
587 + // Attempt reconnection
588 + err := g.AddRelay(addr, dialer)
589 + if err == nil {
590 + log.Info().Str("relay", addr).Msg("[SDK] Reconnection successful")
591 + return
592 + }
593
620 - // Wait before next retry
621 - select {
622 - case <-g.stopch:
594 + if err == ErrRelayExists {
595 + log.Debug().Str("relay", addr).Msg("[SDK] Relay already exists, reconnection complete")
596 return
624 - case <-relay.stop:
597 + }
598 +
599 + retries++
600 +
601 + // Check retry limit (0 or negative means infinite retries)
602 + if maxRetries > 0 && retries >= maxRetries {
603 + log.Error().
604 + Err(err).
605 + Str("relay", addr).
606 + Int("retries", retries).
607 + Msg("[SDK] Reconnection failed after max retries")
608 return
626 - case <-time.After(config.ReconnectInterval):
627 - continue
609 }
629 - }
610
631 - relay.mu.Lock()
632 - relay.client = relayClient
633 - relay.mu.Unlock()
611 + log.Warn().
612 + Err(err).
613 + Str("relay", addr).
614 + Int("attempt", retries).
615 + Msg("[SDK] Reconnection failed, retrying")
616
635 - log.Info().
636 - Str("relay", relay.addr).
637 - Int("attempt", attempt).
638 - Msg("[SDK] Successfully reconnected to relay")
639 -
640 - // Re-register all leases with the reconnected relay
641 - g.mu.Lock()
642 - for _, listener := range g.listeners {
643 - go func(l *RDListener) {
644 - l.mu.Lock()
645 - cred := l.cred
646 - lease := l.lease
647 - l.mu.Unlock()
648 -
649 - err := relayClient.RegisterLease(cred, lease)
650 - if err != nil {
651 - log.Error().
652 - Err(err).
653 - Str("relay", relay.addr).
654 - Str("lease_id", cred.ID()).
655 - Msg("[SDK] Failed to re-register lease after reconnection")
656 - } else {
657 - log.Debug().
658 - Str("relay", relay.addr).
659 - Str("lease_id", cred.ID()).
660 - Msg("[SDK] Lease re-registered after reconnection")
661 - }
662 - }(listener)
617 + // Wait before next retry with context awareness
618 + select {
619 + case <-g.stopch:
620 + log.Debug().Str("relay", addr).Msg("[SDK] Reconnection cancelled during wait")
621 + return
622 + case <-time.After(g.config.ReconnectInterval):
623 + // Continue to next retry
624 + }
625 }
664 - g.mu.Unlock()
665 -
666 - return
667 - }
626 + }()
627 }
628
629 // Implement net.Listener interface for RDListener
@@ -692,8 +651,7 @@ func (l *RDListener) Close() error {
651 // Close all active connections
652 for conn := range l.conns {
653 if err := conn.Close(); err != nil {
695 - // Log error but continue closing other connections
696 - // In a real implementation, you might want to collect errors
654 + log.Error().Err(err).Msg("[SDK] Error closing connection")
655 }
656 delete(l.conns, conn)
657 }
@@ -711,17 +669,16 @@ func (l *RDListener) Addr() net.Addr {
669 // AddRelay adds a new relay server to the client
670 func (g *RDClient) AddRelay(addr string, dialer func(context.Context, string) (io.ReadWriteCloser, error)) error {
671 g.mu.Lock()
672 + defer g.mu.Unlock()
673
674 // Check if relay already exists
675 if _, exists := g.relays[addr]; exists {
717 - g.mu.Unlock()
718 - return errors.New("relay already exists")
676 + return ErrRelayExists
677 }
678
679 // Connect to relay
680 conn, err := dialer(context.Background(), addr)
681 if err != nil {
724 - g.mu.Unlock()
682 return err
683 }
684
@@ -729,8 +686,7 @@ func (g *RDClient) AddRelay(addr string, dialer func(context.Context, string) (i
686 relayClient := portal.NewRelayClient(conn)
687 if relayClient == nil {
688 conn.Close()
732 - g.mu.Unlock()
733 - return errors.New("failed to create relay client")
689 + return ErrRelayNotFound
690 }
691
692 // Add relay
@@ -744,12 +700,9 @@ func (g *RDClient) AddRelay(addr string, dialer func(context.Context, string) (i
700
701 // Register all existing leases with the new relay
702 for _, listener := range g.listeners {
747 - go func(l *RDListener) {
748 - l.mu.Lock()
749 - cred := l.cred
750 - lease := l.lease
751 - l.mu.Unlock()
752 -
703 + cred := listener.cred // immutable
704 + lease := listener.lease.CloneVT()
705 + go func(cred *cryptoops.Credential, lease *rdverb.Lease) {
706 err := relayClient.RegisterLease(cred, lease)
707 if err != nil {
708 log.Error().
@@ -763,18 +716,16 @@ func (g *RDClient) AddRelay(addr string, dialer func(context.Context, string) (i
716 Str("lease_id", cred.ID()).
717 Msg("[SDK] Lease registered with new relay")
718 }
766 - }(listener)
719 + }(cred, lease)
720 }
721
722 // Start listener worker for the new relay
723 g.waitGroup.Add(1)
724 go g.listenerWorker(relay)
725
773 - g.mu.Unlock()
774 -
726 // Start health monitoring for this relay
727 g.waitGroup.Add(1)
777 - go g.healthCheckWorker(relay, g.config)
728 + go g.healthCheckWorker(relay)
729
730 log.Info().Str("relay", addr).Msg("[SDK] New relay added successfully")
731
@@ -784,24 +735,36 @@ func (g *RDClient) AddRelay(addr string, dialer func(context.Context, string) (i
735 // RemoveRelay removes a relay server from the client
736 func (g *RDClient) RemoveRelay(addr string) error {
737 g.mu.Lock()
787 - defer g.mu.Unlock()
788 -
738 relay, exists := g.relays[addr]
739 if !exists {
791 - return errors.New("relay not found")
740 + g.mu.Unlock()
741 + return ErrRelayNotFound
742 }
743
794 - // Signal relay to stop
795 - close(relay.stop)
744 + // Remove from map immediately to prevent duplicate removals
745 + delete(g.relays, addr)
746 + g.mu.Unlock()
747 +
748 + log.Debug().Str("relay", addr).Msg("[SDK] Removing relay")
749 +
750 + // Signal relay to stop (only once)
751 + relay.stopOnce.Do(func() {
752 + close(relay.stop)
753 + })
754
755 // Close relay client
798 - if err := relay.client.Close(); err != nil {
799 - return err
800 - }
756 + relay.mu.Lock()
757 + client := relay.client
758 + relay.mu.Unlock()
759
802 - // Remove from map
803 - delete(g.relays, addr)
760 + if client != nil {
761 + if err := client.Close(); err != nil {
762 + log.Error().Err(err).Str("relay", addr).Msg("[SDK] Error closing relay client")
763 + return err
764 + }
765 + }
766
767 + log.Debug().Str("relay", addr).Msg("[SDK] Relay removed successfully")
768 return nil
769 }
770