mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-09 15:09:08 +02:00
[management] Tear down only the peer session that owns the stream (#8057)
This commit is contained in:
@@ -313,7 +313,7 @@ func (s *Server) Sync(req *proto.EncryptedMessage, srv proto.ManagementService_S
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Debugf("error while sending initial sync for %s: %v", peerKey.String(), err)
|
||||
s.syncSem.Add(-1)
|
||||
s.cancelPeerRoutinesWithoutLock(ctx, accountID, peer, syncStart)
|
||||
s.cancelPeerRoutinesWithoutLock(ctx, accountID, peer, syncStart, nil)
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -321,7 +321,7 @@ func (s *Server) Sync(req *proto.EncryptedMessage, srv proto.ManagementService_S
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Debugf("error while notify peer connected for %s: %v", peerKey.String(), err)
|
||||
s.syncSem.Add(-1)
|
||||
s.cancelPeerRoutinesWithoutLock(ctx, accountID, peer, syncStart)
|
||||
s.cancelPeerRoutinesWithoutLock(ctx, accountID, peer, syncStart, nil)
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -337,7 +337,7 @@ func (s *Server) Sync(req *proto.EncryptedMessage, srv proto.ManagementService_S
|
||||
|
||||
s.syncSem.Add(-1)
|
||||
|
||||
return PeerUpdateHandlerFactory(peerKey, updates, s.secretsManager, srv, func() { s.cancelPeerRoutines(ctx, accountID, peer, syncStart) }).
|
||||
return PeerUpdateHandlerFactory(peerKey, updates, s.secretsManager, srv, func() { s.cancelPeerRoutines(ctx, accountID, peer, syncStart, updates) }).
|
||||
WithMetrics(s.appMetrics).HandleUpdates(ctx)
|
||||
}
|
||||
|
||||
@@ -383,7 +383,7 @@ func (s *Server) startResponseReceiver(ctx context.Context, srv proto.Management
|
||||
|
||||
func (s *Server) sendJobsLoop(ctx context.Context, accountID string, peerKey wgtypes.Key, peer *nbpeer.Peer, updates *job.Channel, srv proto.ManagementService_JobServer) error {
|
||||
// todo figure out better error handling strategy
|
||||
defer s.jobManager.CloseChannel(ctx, accountID, peer.ID)
|
||||
defer s.jobManager.CloseChannel(ctx, accountID, peer.ID, updates)
|
||||
|
||||
for {
|
||||
event, err := updates.Event(ctx)
|
||||
@@ -430,20 +430,26 @@ func (s *Server) sendJob(ctx context.Context, peerKey wgtypes.Key, job *job.Even
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Server) cancelPeerRoutines(ctx context.Context, accountID string, peer *nbpeer.Peer, streamStartTime time.Time) {
|
||||
func (s *Server) cancelPeerRoutines(ctx context.Context, accountID string, peer *nbpeer.Peer, streamStartTime time.Time, session chan *network_map.UpdateMessage) {
|
||||
uncanceledCTX := context.WithoutCancel(ctx)
|
||||
unlock := s.acquirePeerLockByUID(uncanceledCTX, peer.Key)
|
||||
defer unlock()
|
||||
|
||||
s.cancelPeerRoutinesWithoutLock(uncanceledCTX, accountID, peer, streamStartTime)
|
||||
s.cancelPeerRoutinesWithoutLock(uncanceledCTX, accountID, peer, streamStartTime, session)
|
||||
}
|
||||
|
||||
func (s *Server) cancelPeerRoutinesWithoutLock(ctx context.Context, accountID string, peer *nbpeer.Peer, streamStartTime time.Time) {
|
||||
// cancelPeerRoutinesWithoutLock tears down the stream of the session identified by streamStartTime
|
||||
// and its updates channel. A nil session means the stream failed before it registered a channel;
|
||||
// the controller then closes any channel still registered.
|
||||
func (s *Server) cancelPeerRoutinesWithoutLock(ctx context.Context, accountID string, peer *nbpeer.Peer, streamStartTime time.Time, session chan *network_map.UpdateMessage) {
|
||||
err := s.accountManager.OnPeerDisconnected(ctx, accountID, peer.Key, streamStartTime)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("failed to disconnect peer %s properly: %v", peer.Key, err)
|
||||
}
|
||||
s.networkMapController.OnPeerDisconnected(ctx, accountID, peer.ID)
|
||||
if !s.networkMapController.OnPeerDisconnected(ctx, accountID, peer.ID, session) {
|
||||
log.WithContext(ctx).Debugf("skipped peer routines teardown for %s: a newer session owns the peer", peer.Key)
|
||||
return
|
||||
}
|
||||
s.secretsManager.CancelRefresh(peer.ID)
|
||||
|
||||
log.WithContext(ctx).Debugf("peer %s has been disconnected", peer.Key)
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
package grpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
|
||||
"github.com/netbirdio/netbird/management/server/account"
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
)
|
||||
|
||||
func TestCancelPeerRoutines_SessionOwnership(t *testing.T) {
|
||||
peer := &nbpeer.Peer{ID: "peer-1", Key: "peer-key"}
|
||||
streamStart := time.Unix(1700000000, 0)
|
||||
session := make(chan *network_map.UpdateMessage)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
session chan *network_map.UpdateMessage
|
||||
ownsPeer bool
|
||||
cancelRefresh bool
|
||||
}{
|
||||
{name: "owning session tears everything down", session: session, ownsPeer: true, cancelRefresh: true},
|
||||
{name: "stale session keeps the newer session's refresh", session: session, ownsPeer: false, cancelRefresh: false},
|
||||
{name: "failed sync without a channel closes the older session", session: nil, ownsPeer: true, cancelRefresh: true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
accountManager := account.NewMockManager(ctrl)
|
||||
controller := network_map.NewMockController(ctrl)
|
||||
secretsManager := NewMockSecretsManager(ctrl)
|
||||
s := &Server{accountManager: accountManager, networkMapController: controller, secretsManager: secretsManager}
|
||||
|
||||
accountManager.EXPECT().OnPeerDisconnected(gomock.Any(), "account-1", peer.Key, streamStart).Return(nil)
|
||||
controller.EXPECT().OnPeerDisconnected(gomock.Any(), "account-1", peer.ID, tt.session).Return(tt.ownsPeer)
|
||||
if tt.cancelRefresh {
|
||||
secretsManager.EXPECT().CancelRefresh(peer.ID)
|
||||
}
|
||||
|
||||
s.cancelPeerRoutines(context.Background(), "account-1", peer, streamStart, tt.session)
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user