[management] Tear down only the peer session that owns the stream (#8057)

This commit is contained in:
Pascal Fischer
2026-10-08 16:17:23 +02:00
committed by GitHub
parent 515a01dd11
commit 06c4c20010
11 changed files with 340 additions and 24 deletions
@@ -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")
})
}