From 06c4c200108b1499dc2b5a085b13f9febf734db6 Mon Sep 17 00:00:00 2001 From: Pascal Fischer <32096965+pascal-fischer@users.noreply.github.com> Date: Thu, 8 Oct 2026 16:17:23 +0200 Subject: [PATCH] [management] Tear down only the peer session that owns the stream (#8057) --- .../network_map/controller/controller.go | 12 ++- .../controller/controller_disconnect_test.go | 78 +++++++++++++++++++ .../controllers/network_map/interface.go | 5 +- .../controllers/network_map/interface_mock.go | 10 ++- .../controllers/network_map/update_channel.go | 4 + .../update_channel/updatechannel.go | 24 ++++++ .../update_channel/updatechannel_test.go | 55 +++++++++++++ management/internals/shared/grpc/server.go | 22 ++++-- .../shared/grpc/server_disconnect_test.go | 48 ++++++++++++ management/server/job/manager.go | 28 +++++-- management/server/job/manager_test.go | 78 +++++++++++++++++++ 11 files changed, 340 insertions(+), 24 deletions(-) create mode 100644 management/internals/controllers/network_map/controller/controller_disconnect_test.go create mode 100644 management/internals/shared/grpc/server_disconnect_test.go create mode 100644 management/server/job/manager_test.go diff --git a/management/internals/controllers/network_map/controller/controller.go b/management/internals/controllers/network_map/controller/controller.go index b9c27e57e..9b8471c1b 100644 --- a/management/internals/controllers/network_map/controller/controller.go +++ b/management/internals/controllers/network_map/controller/controller.go @@ -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 diff --git a/management/internals/controllers/network_map/controller/controller_disconnect_test.go b/management/internals/controllers/network_map/controller/controller_disconnect_test.go new file mode 100644 index 000000000..ed1a217d9 --- /dev/null +++ b/management/internals/controllers/network_map/controller/controller_disconnect_test.go @@ -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") +} diff --git a/management/internals/controllers/network_map/interface.go b/management/internals/controllers/network_map/interface.go index f447387b4..89345fd58 100644 --- a/management/internals/controllers/network_map/interface.go +++ b/management/internals/controllers/network_map/interface.go @@ -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) } diff --git a/management/internals/controllers/network_map/interface_mock.go b/management/internals/controllers/network_map/interface_mock.go index 5dcd241e1..f79b93019 100644 --- a/management/internals/controllers/network_map/interface_mock.go +++ b/management/internals/controllers/network_map/interface_mock.go @@ -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. diff --git a/management/internals/controllers/network_map/update_channel.go b/management/internals/controllers/network_map/update_channel.go index 0b085b85f..e61432069 100644 --- a/management/internals/controllers/network_map/update_channel.go +++ b/management/internals/controllers/network_map/update_channel.go @@ -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) diff --git a/management/internals/controllers/network_map/update_channel/updatechannel.go b/management/internals/controllers/network_map/update_channel/updatechannel.go index 91627bf15..a8e29afa0 100644 --- a/management/internals/controllers/network_map/update_channel/updatechannel.go +++ b/management/internals/controllers/network_map/update_channel/updatechannel.go @@ -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() diff --git a/management/internals/controllers/network_map/update_channel/updatechannel_test.go b/management/internals/controllers/network_map/update_channel/updatechannel_test.go index c73baf81f..bceefecf3 100644 --- a/management/internals/controllers/network_map/update_channel/updatechannel_test.go +++ b/management/internals/controllers/network_map/update_channel/updatechannel_test.go @@ -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") + }) +} diff --git a/management/internals/shared/grpc/server.go b/management/internals/shared/grpc/server.go index 079a16a70..529e3d379 100644 --- a/management/internals/shared/grpc/server.go +++ b/management/internals/shared/grpc/server.go @@ -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) diff --git a/management/internals/shared/grpc/server_disconnect_test.go b/management/internals/shared/grpc/server_disconnect_test.go new file mode 100644 index 000000000..9a19a9b1a --- /dev/null +++ b/management/internals/shared/grpc/server_disconnect_test.go @@ -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) + }) + } +} diff --git a/management/server/job/manager.go b/management/server/job/manager.go index 0b183ac39..323d1e639 100644 --- a/management/server/job/manager.go +++ b/management/server/job/manager.go @@ -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) } } diff --git a/management/server/job/manager_test.go b/management/server/job/manager_test.go new file mode 100644 index 000000000..4919ae869 --- /dev/null +++ b/management/server/job/manager_test.go @@ -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") +}