From 7d8e20030b2e068b9585fdb65c6b4c90931719ed Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Mon, 29 Jun 2026 00:41:47 +0200 Subject: [PATCH] [relay] Extract foreign relay client cache into a dedicated type Move the foreign-relay client cache out of Manager into a foreignRelays type. Concurrent first-time connects to the same server are deduplicated with singleflight, so the cache mutex is never held during a network connect (removing the previous stall where a slow connect blocked all map operations). A per-entry in-use refcount prevents the cleanup loop from closing a client while a connection is being opened on it. This drops RelayTrack, its per-track lock and the hand-over-hand locking between the map lock and the track lock. The exported API is unchanged. --- shared/relay/client/foreign_relays.go | 156 ++++++++++++++++++++++++++ shared/relay/client/manager.go | 147 +----------------------- 2 files changed, 162 insertions(+), 141 deletions(-) create mode 100644 shared/relay/client/foreign_relays.go diff --git a/shared/relay/client/foreign_relays.go b/shared/relay/client/foreign_relays.go new file mode 100644 index 000000000..fed510c5f --- /dev/null +++ b/shared/relay/client/foreign_relays.go @@ -0,0 +1,156 @@ +package client + +import ( + "context" + "net" + "net/netip" + "sync" + "time" + + log "github.com/sirupsen/logrus" + "golang.org/x/sync/singleflight" + + relayAuth "github.com/netbirdio/netbird/shared/relay/auth/hmac" +) + +type foreignRelay struct { + client *Client + created time.Time + inUse int +} + +type foreignRelays struct { + mu sync.RWMutex + clients map[string]*foreignRelay + + group singleflight.Group + + ctx context.Context + tokenStore *relayAuth.TokenStore + peerID string + mtu uint16 + transportFallback *transportFallback + onDisconnect func(string) + keepUnusedServerTime time.Duration +} + +func newForeignRelays(ctx context.Context, tokenStore *relayAuth.TokenStore, peerID string, mtu uint16, transportFallback *transportFallback, onDisconnect func(string), keepUnusedServerTime time.Duration) *foreignRelays { + return &foreignRelays{ + clients: make(map[string]*foreignRelay), + ctx: ctx, + tokenStore: tokenStore, + peerID: peerID, + mtu: mtu, + transportFallback: transportFallback, + onDisconnect: onDisconnect, + keepUnusedServerTime: keepUnusedServerTime, + } +} + +func (f *foreignRelays) openConn(ctx context.Context, serverAddress, peerKey string, serverIP netip.Addr) (net.Conn, error) { + fr, err := f.acquire(serverAddress, serverIP) + if err != nil { + return nil, err + } + defer f.release(fr) + + return fr.client.OpenConn(ctx, peerKey) +} + +func (f *foreignRelays) acquire(serverAddress string, serverIP netip.Addr) (*foreignRelay, error) { + f.mu.Lock() + if fr, ok := f.clients[serverAddress]; ok { + fr.inUse++ + f.mu.Unlock() + return fr, nil + } + f.mu.Unlock() + + v, err, _ := f.group.Do(serverAddress, func() (any, error) { + f.mu.RLock() + fr, ok := f.clients[serverAddress] + f.mu.RUnlock() + if ok { + return fr, nil + } + + relayClient := NewClientWithServerIP(serverAddress, serverIP, f.tokenStore, f.peerID, f.mtu) + relayClient.SetTransportFallback(f.transportFallback) + if err := relayClient.Connect(f.ctx); err != nil { + return nil, err + } + relayClient.SetOnDisconnectListener(f.onDisconnect) + + f.mu.Lock() + fr = &foreignRelay{client: relayClient, created: time.Now()} + f.clients[serverAddress] = fr + f.mu.Unlock() + return fr, nil + }) + if err != nil { + return nil, err + } + + fr := v.(*foreignRelay) + f.mu.Lock() + if cur, ok := f.clients[serverAddress]; !ok || cur != fr { + f.mu.Unlock() + return f.acquire(serverAddress, serverIP) + } + fr.inUse++ + f.mu.Unlock() + return fr, nil +} + +func (f *foreignRelays) release(fr *foreignRelay) { + f.mu.Lock() + fr.inUse-- + f.mu.Unlock() +} + +func (f *foreignRelays) evict(serverAddress string) { + f.mu.Lock() + defer f.mu.Unlock() + if _, ok := f.clients[serverAddress]; ok { + delete(f.clients, serverAddress) + log.Debugf("evicted disconnected foreign relay client: %s", serverAddress) + } +} + +func (f *foreignRelays) cleanupUnused() { + f.mu.Lock() + defer f.mu.Unlock() + + for addr, fr := range f.clients { + if time.Since(fr.created) <= f.keepUnusedServerTime { + continue + } + if fr.inUse > 0 { + continue + } + if fr.client.HasConns() { + continue + } + fr.client.SetOnDisconnectListener(nil) + go func() { + _ = fr.client.Close() + }() + log.Debugf("clean up unused relay server connection: %s", addr) + delete(f.clients, addr) + } +} + +func (f *foreignRelays) states() []RelayConnState { + f.mu.RLock() + clients := make([]*Client, 0, len(f.clients)) + for _, fr := range f.clients { + clients = append(clients, fr.client) + } + f.mu.RUnlock() + + states := make([]RelayConnState, 0, len(clients)) + for _, c := range clients { + states = append(states, relayConnState(c)) + } + return states +} diff --git a/shared/relay/client/manager.go b/shared/relay/client/manager.go index e1515401e..05cf2fa26 100644 --- a/shared/relay/client/manager.go +++ b/shared/relay/client/manager.go @@ -22,22 +22,6 @@ var ( ErrRelayClientNotConnected = fmt.Errorf("relay client not connected") ) -// RelayTrack hold the relay clients for the foreign relay servers. -// With the mutex can ensure we can open new connection in case the relay connection has been established with -// the relay server. -type RelayTrack struct { - sync.RWMutex - relayClient *Client - err error - created time.Time -} - -func NewRelayTrack() *RelayTrack { - return &RelayTrack{ - created: time.Now(), - } -} - type OnServerCloseListener func() // ManagerOption configures a Manager at construction time. @@ -78,8 +62,7 @@ type Manager struct { relayClientMu sync.RWMutex reconnectGuard *Guard - relayClients map[string]*RelayTrack - relayClientsMutex sync.RWMutex + foreign *foreignRelays onDisconnectedListeners map[string]*list.List onReconnectedListenerFn func() @@ -115,7 +98,6 @@ func NewManager(ctx context.Context, serverURLs []string, peerID string, mtu uin ConnectionTimeout: defaultConnectionTimeout, TransportFallback: tf, }, - relayClients: make(map[string]*RelayTrack), onDisconnectedListeners: make(map[string]*list.List), cleanupInterval: relayCleanupInterval, keepUnusedServerTime: keepUnusedServerTime, @@ -123,6 +105,7 @@ func NewManager(ctx context.Context, serverURLs []string, peerID string, mtu uin for _, opt := range opts { opt(m) } + m.foreign = newForeignRelays(ctx, tokenStore, peerID, mtu, tf, m.onServerDisconnected, m.keepUnusedServerTime) m.serverPicker.ServerURLs.Store(serverURLs) m.reconnectGuard = NewGuard(m.serverPicker, m.maxBackoffInterval) return m @@ -181,7 +164,7 @@ func (m *Manager) OpenConn(ctx context.Context, serverAddress, peerKey string, s netConn, err = m.relayClient.OpenConn(ctx, peerKey) } else { log.Debugf("open peer connection via foreign server: %s", serverAddress) - netConn, err = m.openConnVia(ctx, serverAddress, peerKey, serverIP) + netConn, err = m.foreign.openConn(ctx, serverAddress, peerKey, serverIP) } if err != nil { return nil, err @@ -282,26 +265,7 @@ func (m *Manager) RelayStates() []RelayConnState { states = append(states, st) } - // Snapshot the tracks, then query each outside the map lock: a track can be - // held by an in-progress Connect, and blocking on it must not stall other - // relay operations. - m.relayClientsMutex.RLock() - tracks := make([]*RelayTrack, 0, len(m.relayClients)) - for _, rt := range m.relayClients { - tracks = append(tracks, rt) - } - m.relayClientsMutex.RUnlock() - - // Only connected foreign relays carry state; a failed connect is evicted - // immediately (openConnVia), so there is no error state to surface. - for _, rt := range tracks { - rt.RLock() - rc := rt.relayClient - rt.RUnlock() - if rc != nil { - states = append(states, relayConnState(rc)) - } - } + states = append(states, m.foreign.states()...) return states } @@ -322,64 +286,6 @@ func (m *Manager) UpdateToken(token *relayAuth.Token) error { return m.tokenStore.UpdateToken(token) } -func (m *Manager) openConnVia(ctx context.Context, serverAddress, peerKey string, serverIP netip.Addr) (net.Conn, error) { - // check if already has a connection to the desired relay server - m.relayClientsMutex.RLock() - rt, ok := m.relayClients[serverAddress] - if ok { - rt.RLock() - m.relayClientsMutex.RUnlock() - defer rt.RUnlock() - if rt.err != nil { - return nil, rt.err - } - return rt.relayClient.OpenConn(ctx, peerKey) - } - m.relayClientsMutex.RUnlock() - - // if not, establish a new connection but check it again (because changed the lock type) before starting the - // connection - m.relayClientsMutex.Lock() - rt, ok = m.relayClients[serverAddress] - if ok { - rt.RLock() - m.relayClientsMutex.Unlock() - defer rt.RUnlock() - if rt.err != nil { - return nil, rt.err - } - return rt.relayClient.OpenConn(ctx, peerKey) - } - - // create a new relay client and store it in the relayClients map - rt = NewRelayTrack() - rt.Lock() - m.relayClients[serverAddress] = rt - m.relayClientsMutex.Unlock() - - relayClient := NewClientWithServerIP(serverAddress, serverIP, m.tokenStore, m.peerID, m.mtu) - relayClient.SetTransportFallback(m.transportFallback) - err := relayClient.Connect(m.ctx) - if err != nil { - rt.err = err - rt.Unlock() - m.relayClientsMutex.Lock() - delete(m.relayClients, serverAddress) - m.relayClientsMutex.Unlock() - return nil, err - } - // if connection closed then delete the relay client from the list - relayClient.SetOnDisconnectListener(m.onServerDisconnected) - rt.relayClient = relayClient - rt.Unlock() - - conn, err := relayClient.OpenConn(ctx, peerKey) - if err != nil { - return nil, err - } - return conn, nil -} - func (m *Manager) onServerConnected() { m.listenerLock.Lock() defer m.listenerLock.Unlock() @@ -405,21 +311,12 @@ func (m *Manager) onServerDisconnected(serverAddress string) { m.relayClientMu.Unlock() if !isHome { - m.evictForeignRelay(serverAddress) + m.foreign.evict(serverAddress) } m.notifyOnDisconnectListeners(serverAddress) } -func (m *Manager) evictForeignRelay(serverAddress string) { - m.relayClientsMutex.Lock() - defer m.relayClientsMutex.Unlock() - if _, ok := m.relayClients[serverAddress]; ok { - delete(m.relayClients, serverAddress) - log.Debugf("evicted disconnected foreign relay client: %s", serverAddress) - } -} - func (m *Manager) listenGuardEvent(ctx context.Context) { for { select { @@ -458,43 +355,11 @@ func (m *Manager) startCleanupLoop() { case <-m.ctx.Done(): return case <-ticker.C: - m.cleanUpUnusedRelays() + m.foreign.cleanupUnused() } } } -func (m *Manager) cleanUpUnusedRelays() { - m.relayClientsMutex.Lock() - defer m.relayClientsMutex.Unlock() - - for addr, rt := range m.relayClients { - rt.Lock() - // if the connection failed to the server the relay client will be nil - // but the instance will be kept in the relayClients until the next locking - if rt.err != nil { - rt.Unlock() - continue - } - - if time.Since(rt.created) <= m.keepUnusedServerTime { - rt.Unlock() - continue - } - - if rt.relayClient.HasConns() { - rt.Unlock() - continue - } - rt.relayClient.SetOnDisconnectListener(nil) - go func() { - _ = rt.relayClient.Close() - }() - log.Debugf("clean up unused relay server connection: %s", addr) - delete(m.relayClients, addr) - rt.Unlock() - } -} - func (m *Manager) addListener(serverAddress string, onClosedListener OnServerCloseListener) { m.listenerLock.Lock() defer m.listenerLock.Unlock()