mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-09 06:59:08 +02:00
[management] Tear down only the peer session that owns the stream (#8057)
This commit is contained in:
@@ -127,14 +127,20 @@ func (c *Controller) OnPeerConnected(ctx context.Context, accountID string, peer
|
||||
return c.peersUpdateManager.CreateChannel(ctx, peerID), nil
|
||||
}
|
||||
|
||||
func (c *Controller) OnPeerDisconnected(ctx context.Context, accountID string, peerID string) {
|
||||
c.peersUpdateManager.CloseChannel(ctx, peerID)
|
||||
// OnPeerDisconnected closes the session's updates channel and schedules an ephemeral peer for
|
||||
// cleanup. It returns false without touching anything when a newer session owns the peer. A nil
|
||||
// session closes any registered channel.
|
||||
func (c *Controller) OnPeerDisconnected(ctx context.Context, accountID string, peerID string, session chan *network_map.UpdateMessage) bool {
|
||||
if !c.peersUpdateManager.CloseSessionChannel(ctx, peerID, session) {
|
||||
return false
|
||||
}
|
||||
peer, err := c.repo.GetPeerByID(ctx, accountID, peerID)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("failed to get peer %s: %v", peerID, err)
|
||||
return
|
||||
return true
|
||||
}
|
||||
c.EphemeralPeersManager.OnPeerDisconnected(ctx, peer)
|
||||
return true
|
||||
}
|
||||
|
||||
// injectAllProxyPolicies prepares an account for the per-peer network-map
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map/update_channel"
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
)
|
||||
|
||||
type recordingEphemeralManager struct {
|
||||
disconnected []string
|
||||
}
|
||||
|
||||
func (r *recordingEphemeralManager) LoadInitialPeers(context.Context) {}
|
||||
func (r *recordingEphemeralManager) Stop() {}
|
||||
func (r *recordingEphemeralManager) OnPeerConnected(context.Context, *nbpeer.Peer) {}
|
||||
func (r *recordingEphemeralManager) OnPeerDisconnected(_ context.Context, peer *nbpeer.Peer) {
|
||||
r.disconnected = append(r.disconnected, peer.ID)
|
||||
}
|
||||
|
||||
func TestOnPeerDisconnected_OwnSession(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
repo := NewMockRepository(gomock.NewController(t))
|
||||
ephemeralManager := &recordingEphemeralManager{}
|
||||
updateManager := update_channel.NewPeersUpdateManager(nil)
|
||||
c := Controller{repo: repo, peersUpdateManager: updateManager, EphemeralPeersManager: ephemeralManager}
|
||||
|
||||
session := updateManager.CreateChannel(ctx, "peer-1")
|
||||
repo.EXPECT().GetPeerByID(gomock.Any(), "account-1", "peer-1").Return(&nbpeer.Peer{ID: "peer-1", Ephemeral: true}, nil)
|
||||
|
||||
require.True(t, c.OnPeerDisconnected(ctx, "account-1", "peer-1", session))
|
||||
assert.False(t, updateManager.HasChannel("peer-1"))
|
||||
assert.Equal(t, []string{"peer-1"}, ephemeralManager.disconnected)
|
||||
}
|
||||
|
||||
func TestOnPeerDisconnected_NilSessionClosesOlderChannel(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
repo := NewMockRepository(gomock.NewController(t))
|
||||
ephemeralManager := &recordingEphemeralManager{}
|
||||
updateManager := update_channel.NewPeersUpdateManager(nil)
|
||||
c := Controller{repo: repo, peersUpdateManager: updateManager, EphemeralPeersManager: ephemeralManager}
|
||||
|
||||
older := updateManager.CreateChannel(ctx, "peer-1")
|
||||
repo.EXPECT().GetPeerByID(gomock.Any(), "account-1", "peer-1").Return(&nbpeer.Peer{ID: "peer-1", Ephemeral: true}, nil)
|
||||
|
||||
require.True(t, c.OnPeerDisconnected(ctx, "account-1", "peer-1", nil))
|
||||
assert.False(t, updateManager.HasChannel("peer-1"))
|
||||
_, open := <-older
|
||||
assert.False(t, open, "older channel must be closed")
|
||||
assert.Equal(t, []string{"peer-1"}, ephemeralManager.disconnected)
|
||||
}
|
||||
|
||||
func TestOnPeerDisconnected_NewerSessionOwnsPeer(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ephemeralManager := &recordingEphemeralManager{}
|
||||
updateManager := update_channel.NewPeersUpdateManager(nil)
|
||||
c := Controller{repo: NewMockRepository(gomock.NewController(t)), peersUpdateManager: updateManager, EphemeralPeersManager: ephemeralManager}
|
||||
|
||||
stale := updateManager.CreateChannel(ctx, "peer-1")
|
||||
current := updateManager.CreateChannel(ctx, "peer-1")
|
||||
|
||||
require.False(t, c.OnPeerDisconnected(ctx, "account-1", "peer-1", stale))
|
||||
require.True(t, updateManager.HasChannel("peer-1"))
|
||||
updateManager.SendUpdate(ctx, "peer-1", &network_map.UpdateMessage{})
|
||||
select {
|
||||
case _, open := <-current:
|
||||
assert.True(t, open, "newer session channel must stay open")
|
||||
default:
|
||||
t.Fatal("newer session channel did not receive the update")
|
||||
}
|
||||
assert.Empty(t, ephemeralManager.disconnected, "a live peer must not be scheduled for ephemeral cleanup")
|
||||
}
|
||||
@@ -35,7 +35,10 @@ type Controller interface {
|
||||
OnPeersDeleted(ctx context.Context, accountID string, peerIDs []string, affectedPeerIDs []string) error
|
||||
DisconnectPeers(ctx context.Context, accountId string, peerIDs []string)
|
||||
OnPeerConnected(ctx context.Context, accountID string, peerID string) (chan *UpdateMessage, error)
|
||||
OnPeerDisconnected(ctx context.Context, accountID string, peerID string)
|
||||
// OnPeerDisconnected tears down the stream state of the peer's session, identified by its
|
||||
// updates channel. It returns false and leaves everything untouched when a newer session
|
||||
// owns the peer. A nil session is the newest session and tears down any registered channel.
|
||||
OnPeerDisconnected(ctx context.Context, accountID string, peerID string, session chan *UpdateMessage) bool
|
||||
|
||||
TrackEphemeralPeer(ctx context.Context, peer *nbpeer.Peer)
|
||||
}
|
||||
|
||||
@@ -177,15 +177,17 @@ func (mr *MockControllerMockRecorder) OnPeerConnected(ctx, accountID, peerID any
|
||||
}
|
||||
|
||||
// OnPeerDisconnected mocks base method.
|
||||
func (m *MockController) OnPeerDisconnected(ctx context.Context, accountID, peerID string) {
|
||||
func (m *MockController) OnPeerDisconnected(ctx context.Context, accountID, peerID string, session chan *UpdateMessage) bool {
|
||||
m.ctrl.T.Helper()
|
||||
m.ctrl.Call(m, "OnPeerDisconnected", ctx, accountID, peerID)
|
||||
ret := m.ctrl.Call(m, "OnPeerDisconnected", ctx, accountID, peerID, session)
|
||||
ret0, _ := ret[0].(bool)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// OnPeerDisconnected indicates an expected call of OnPeerDisconnected.
|
||||
func (mr *MockControllerMockRecorder) OnPeerDisconnected(ctx, accountID, peerID any) *gomock.Call {
|
||||
func (mr *MockControllerMockRecorder) OnPeerDisconnected(ctx, accountID, peerID, session any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "OnPeerDisconnected", reflect.TypeOf((*MockController)(nil).OnPeerDisconnected), ctx, accountID, peerID)
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "OnPeerDisconnected", reflect.TypeOf((*MockController)(nil).OnPeerDisconnected), ctx, accountID, peerID, session)
|
||||
}
|
||||
|
||||
// OnPeersAdded mocks base method.
|
||||
|
||||
@@ -6,6 +6,10 @@ type PeersUpdateManager interface {
|
||||
SendUpdate(ctx context.Context, peerID string, update *UpdateMessage)
|
||||
CreateChannel(ctx context.Context, peerID string) chan *UpdateMessage
|
||||
CloseChannel(ctx context.Context, peerID string)
|
||||
// CloseSessionChannel closes the peer's channel and returns true, unless a channel of
|
||||
// another session is registered, in which case it returns false and closes nothing.
|
||||
// A nil session is the newest session and closes any registered channel.
|
||||
CloseSessionChannel(ctx context.Context, peerID string, session chan *UpdateMessage) bool
|
||||
CountStreams() int
|
||||
HasChannel(peerID string) bool
|
||||
CloseChannels(ctx context.Context, peerIDs []string)
|
||||
|
||||
@@ -133,6 +133,30 @@ func (p *PeersUpdateManager) CloseChannel(ctx context.Context, peerID string) {
|
||||
p.closeChannel(ctx, peerID)
|
||||
}
|
||||
|
||||
// CloseSessionChannel closes the peer's updates channel only while it is still session's channel.
|
||||
// It returns false when a newer stream has registered a different channel, which the stale
|
||||
// session must leave to its new owner. A nil session belongs to a stream that failed before
|
||||
// registering its own channel; it is the newest session and closes any channel still registered.
|
||||
func (p *PeersUpdateManager) CloseSessionChannel(ctx context.Context, peerID string, session chan *network_map.UpdateMessage) bool {
|
||||
start := time.Now()
|
||||
|
||||
p.channelsMux.Lock()
|
||||
defer func() {
|
||||
p.channelsMux.Unlock()
|
||||
if p.metrics != nil {
|
||||
p.metrics.UpdateChannelMetrics().CountCloseChannelDuration(time.Since(start))
|
||||
}
|
||||
}()
|
||||
|
||||
if channel, ok := p.peerChannels[peerID]; ok && session != nil && channel != session {
|
||||
log.WithContext(ctx).Debugf("skipped closing updates channel: peer %s is owned by a newer session", peerID)
|
||||
return false
|
||||
}
|
||||
|
||||
p.closeChannel(ctx, peerID)
|
||||
return true
|
||||
}
|
||||
|
||||
// GetAllConnectedPeers returns a copy of the connected peers map
|
||||
func (p *PeersUpdateManager) GetAllConnectedPeers() map[string]struct{} {
|
||||
start := time.Now()
|
||||
|
||||
@@ -5,6 +5,9 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
@@ -84,3 +87,55 @@ func TestCloseChannel(t *testing.T) {
|
||||
t.Error("Error closing the channel")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloseSessionChannel(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
const peer = "test-close-session"
|
||||
|
||||
t.Run("own channel is closed", func(t *testing.T) {
|
||||
peersUpdater := NewPeersUpdateManager(nil)
|
||||
session := peersUpdater.CreateChannel(ctx, peer)
|
||||
|
||||
require.True(t, peersUpdater.CloseSessionChannel(ctx, peer, session))
|
||||
assert.False(t, peersUpdater.HasChannel(peer))
|
||||
_, open := <-session
|
||||
assert.False(t, open, "own channel must be closed")
|
||||
})
|
||||
|
||||
t.Run("newer session channel is kept", func(t *testing.T) {
|
||||
peersUpdater := NewPeersUpdateManager(nil)
|
||||
stale := peersUpdater.CreateChannel(ctx, peer)
|
||||
current := peersUpdater.CreateChannel(ctx, peer)
|
||||
|
||||
require.False(t, peersUpdater.CloseSessionChannel(ctx, peer, stale))
|
||||
require.True(t, peersUpdater.HasChannel(peer))
|
||||
assert.Equal(t, current, peersUpdater.peerChannels[peer])
|
||||
|
||||
peersUpdater.SendUpdate(ctx, peer, &network_map.UpdateMessage{})
|
||||
select {
|
||||
case _, open := <-current:
|
||||
assert.True(t, open, "newer session channel must stay open")
|
||||
default:
|
||||
t.Fatal("newer session channel did not receive the update")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("no registered channel", func(t *testing.T) {
|
||||
peersUpdater := NewPeersUpdateManager(nil)
|
||||
session := peersUpdater.CreateChannel(ctx, peer)
|
||||
peersUpdater.CloseChannel(ctx, peer)
|
||||
|
||||
assert.True(t, peersUpdater.CloseSessionChannel(ctx, peer, session))
|
||||
assert.True(t, peersUpdater.CloseSessionChannel(ctx, peer, nil))
|
||||
})
|
||||
|
||||
t.Run("nil session closes the registered channel", func(t *testing.T) {
|
||||
peersUpdater := NewPeersUpdateManager(nil)
|
||||
current := peersUpdater.CreateChannel(ctx, peer)
|
||||
|
||||
require.True(t, peersUpdater.CloseSessionChannel(ctx, peer, nil))
|
||||
assert.False(t, peersUpdater.HasChannel(peer))
|
||||
_, open := <-current
|
||||
assert.False(t, open, "registered channel must be closed")
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -57,6 +57,7 @@ func (jm *Manager) CreateJobChannel(ctx context.Context, accountID, peerID strin
|
||||
if ch, ok := jm.jobChannels[peerID]; ok {
|
||||
ch.Close()
|
||||
delete(jm.jobChannels, peerID)
|
||||
jm.failPendingLocked(ctx, accountID, peerID, "Pending job cleanup: job stream replaced by a newer stream")
|
||||
}
|
||||
|
||||
ch := NewChannel()
|
||||
@@ -127,24 +128,35 @@ func (jm *Manager) HandleResponse(ctx context.Context, resp *proto.JobResponse,
|
||||
return nil
|
||||
}
|
||||
|
||||
// CloseChannel closes a peer’s channel and cleans up its jobs
|
||||
func (jm *Manager) CloseChannel(ctx context.Context, accountID, peerID string) {
|
||||
// CloseChannel closes the peer's job channel session and fails its pending jobs. It does nothing
|
||||
// when a newer job stream has registered a different channel for the peer.
|
||||
func (jm *Manager) CloseChannel(ctx context.Context, accountID, peerID string, session *Channel) {
|
||||
jm.mu.Lock()
|
||||
defer jm.mu.Unlock()
|
||||
|
||||
if ch, ok := jm.jobChannels[peerID]; ok {
|
||||
if ch != session {
|
||||
log.WithContext(ctx).Debugf("skipped closing job channel: peer %s is owned by a newer job stream", peerID)
|
||||
return
|
||||
}
|
||||
ch.Close()
|
||||
delete(jm.jobChannels, peerID)
|
||||
}
|
||||
|
||||
jm.failPendingLocked(ctx, accountID, peerID, "Time out peer disconnected")
|
||||
}
|
||||
|
||||
// failPendingLocked marks the peer's pending jobs as failed and drops them from memory. The caller
|
||||
// must hold jm.mu.
|
||||
func (jm *Manager) failPendingLocked(ctx context.Context, accountID, peerID, reason string) {
|
||||
for jobID, ev := range jm.pending {
|
||||
if ev.PeerID == peerID {
|
||||
// if the client disconnect and there is pending job then mark it as failed
|
||||
if err := jm.Store.MarkPendingJobsAsFailed(ctx, accountID, peerID, jobID, "Time out peer disconnected"); err != nil {
|
||||
log.WithContext(ctx).Errorf("failed to mark pending jobs as failed: %v", err)
|
||||
}
|
||||
delete(jm.pending, jobID)
|
||||
if ev.PeerID != peerID {
|
||||
continue
|
||||
}
|
||||
if err := jm.Store.MarkPendingJobsAsFailed(ctx, accountID, peerID, jobID, reason); err != nil {
|
||||
log.WithContext(ctx).Errorf("failed to mark pending jobs as failed: %v", err)
|
||||
}
|
||||
delete(jm.pending, jobID)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
package job
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
)
|
||||
|
||||
func TestCloseChannel_OwnSession(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mockStore := store.NewMockStore(gomock.NewController(t))
|
||||
jm := NewJobManager(nil, mockStore, nil)
|
||||
|
||||
session := NewChannel()
|
||||
jm.jobChannels["peer-1"] = session
|
||||
jm.pending["job-1"] = &Event{PeerID: "peer-1"}
|
||||
jm.pending["job-2"] = &Event{PeerID: "peer-2"}
|
||||
|
||||
mockStore.EXPECT().MarkPendingJobsAsFailed(gomock.Any(), "account-1", "peer-1", "job-1", gomock.Any()).Return(nil)
|
||||
|
||||
jm.CloseChannel(ctx, "account-1", "peer-1", session)
|
||||
|
||||
assert.False(t, jm.IsPeerConnected("peer-1"))
|
||||
_, err := session.Event(ctx)
|
||||
assert.ErrorIs(t, err, ErrJobChannelClosed)
|
||||
assert.NotContains(t, jm.pending, "job-1")
|
||||
assert.Contains(t, jm.pending, "job-2", "jobs of other peers must be kept")
|
||||
}
|
||||
|
||||
func TestCloseChannel_NewerSessionOwnsPeer(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
jm := NewJobManager(nil, store.NewMockStore(gomock.NewController(t)), nil)
|
||||
|
||||
stale := NewChannel()
|
||||
current := NewChannel()
|
||||
jm.jobChannels["peer-1"] = current
|
||||
jm.pending["job-1"] = &Event{PeerID: "peer-1"}
|
||||
|
||||
jm.CloseChannel(ctx, "account-1", "peer-1", stale)
|
||||
|
||||
require.True(t, jm.IsPeerConnected("peer-1"))
|
||||
assert.Contains(t, jm.pending, "job-1", "jobs served by the newer stream must not be failed")
|
||||
|
||||
event := &Event{PeerID: "peer-1"}
|
||||
require.NoError(t, current.AddEvent(ctx, time.Second, event))
|
||||
got, err := current.Event(ctx)
|
||||
require.NoError(t, err, "newer session channel must stay open")
|
||||
assert.Equal(t, event, got)
|
||||
}
|
||||
|
||||
func TestCreateJobChannel_ReplacesSessionAndDropsItsPendingJobs(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mockStore := store.NewMockStore(gomock.NewController(t))
|
||||
jm := NewJobManager(nil, mockStore, nil)
|
||||
|
||||
stale := NewChannel()
|
||||
jm.jobChannels["peer-1"] = stale
|
||||
jm.pending["job-1"] = &Event{PeerID: "peer-1"}
|
||||
jm.pending["job-2"] = &Event{PeerID: "peer-2"}
|
||||
|
||||
mockStore.EXPECT().MarkAllPendingJobsAsFailed(gomock.Any(), "account-1", "peer-1", gomock.Any()).Return(nil)
|
||||
mockStore.EXPECT().MarkPendingJobsAsFailed(gomock.Any(), "account-1", "peer-1", "job-1", "Pending job cleanup: job stream replaced by a newer stream").Return(nil)
|
||||
|
||||
current := jm.CreateJobChannel(ctx, "account-1", "peer-1")
|
||||
|
||||
require.NotSame(t, stale, current)
|
||||
_, err := stale.Event(ctx)
|
||||
assert.ErrorIs(t, err, ErrJobChannelClosed, "replaced channel must be closed")
|
||||
assert.True(t, jm.IsPeerConnected("peer-1"))
|
||||
assert.False(t, jm.IsPeerHasPendingJobs("peer-1"), "jobs of the replaced stream must be dropped")
|
||||
assert.Contains(t, jm.pending, "job-2", "jobs of other peers must be kept")
|
||||
}
|
||||
Reference in New Issue
Block a user