diff --git a/management/internals/controllers/network_map/controller/controller.go b/management/internals/controllers/network_map/controller/controller.go index 07f1938c5..30de974a1 100644 --- a/management/internals/controllers/network_map/controller/controller.go +++ b/management/internals/controllers/network_map/controller/controller.go @@ -651,6 +651,11 @@ func (c *Controller) GetValidatedPeerWithComponents(ctx context.Context, isRequi return nil, nil, nil, nil, 0, err } + // it's possible that the peer gets deleted between the call to "sendInitialSync()" and here, bail out in this case + if _, ok := account.Peers[peer.ID]; !ok { + return nil, nil, nil, nil, 0, fmt.Errorf("peer '%s' no longer exists", peer.ID) + } + c.injectAllProxyPolicies(ctx, account) approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra) diff --git a/management/internals/controllers/network_map/controller/controller_test.go b/management/internals/controllers/network_map/controller/controller_test.go index 90e7b6e18..dfbbb2915 100644 --- a/management/internals/controllers/network_map/controller/controller_test.go +++ b/management/internals/controllers/network_map/controller/controller_test.go @@ -1,10 +1,15 @@ package controller import ( + "context" "testing" "github.com/netbirdio/netbird/management/internals/controllers/network_map" + "github.com/netbirdio/netbird/management/server/account" nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/management/server/types" + "github.com/stretchr/testify/assert" + "go.uber.org/mock/gomock" ) func TestComputeForwarderPort(t *testing.T) { @@ -107,3 +112,22 @@ func TestComputeForwarderPort(t *testing.T) { t.Errorf("Expected %d for peers with unknown version, got %d", network_map.OldForwarderPort, result) } } + +func TestGetValidatedPeerWithComponents_DeletedPeer(t *testing.T) { + ctrl := gomock.NewController(t) + mockrequestBuffer := account.NewMockRequestBuffer(ctrl) + + c := Controller{ + requestBuffer: mockrequestBuffer, + } + + mockrequestBuffer.EXPECT().GetAccountWithBackpressure(gomock.Any(), gomock.Any()).Return(&types.Account{}, nil) + peer, components, netmap, posturechecks, dnsforwardPort, err := c.GetValidatedPeerWithComponents(context.TODO(), false, "test-account-id", &nbpeer.Peer{ID: "test-peer-id"}) + + assert.Nil(t, peer) + assert.Nil(t, components) + assert.Nil(t, netmap) + assert.Nil(t, posturechecks) + assert.Equal(t, int64(0), dnsforwardPort) + assert.NotNil(t, err) +} diff --git a/management/internals/controllers/network_map/controller/repository.go b/management/internals/controllers/network_map/controller/repository.go index c0fcefc7d..bd8ed4e80 100644 --- a/management/internals/controllers/network_map/controller/repository.go +++ b/management/internals/controllers/network_map/controller/repository.go @@ -3,14 +3,16 @@ package controller import ( "context" + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" "github.com/netbirdio/netbird/management/internals/modules/zones" - "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" "github.com/netbirdio/netbird/management/server/peer" "github.com/netbirdio/netbird/management/server/store" "github.com/netbirdio/netbird/management/server/types" ) +//go:generate go tool mockgen -source=./repository.go -package=controller -destination=repository_mock.go + type Repository interface { GetAccountNetwork(ctx context.Context, accountID string) (*types.Network, error) GetAccountPeers(ctx context.Context, accountID string) ([]*peer.Peer, error) diff --git a/management/internals/controllers/network_map/controller/repository_mock.go b/management/internals/controllers/network_map/controller/repository_mock.go new file mode 100644 index 000000000..5246eef4b --- /dev/null +++ b/management/internals/controllers/network_map/controller/repository_mock.go @@ -0,0 +1,150 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: ./repository.go +// +// Generated by this command: +// +// mockgen -source=./repository.go -package=controller -destination=repository_mock.go +// + +// Package controller is a generated GoMock package. +package controller + +import ( + context "context" + reflect "reflect" + + service "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + zones "github.com/netbirdio/netbird/management/internals/modules/zones" + peer "github.com/netbirdio/netbird/management/server/peer" + types "github.com/netbirdio/netbird/management/server/types" + gomock "go.uber.org/mock/gomock" +) + +// MockRepository is a mock of Repository interface. +type MockRepository struct { + ctrl *gomock.Controller + recorder *MockRepositoryMockRecorder + isgomock struct{} +} + +// MockRepositoryMockRecorder is the mock recorder for MockRepository. +type MockRepositoryMockRecorder struct { + mock *MockRepository +} + +// NewMockRepository creates a new mock instance. +func NewMockRepository(ctrl *gomock.Controller) *MockRepository { + mock := &MockRepository{ctrl: ctrl} + mock.recorder = &MockRepositoryMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockRepository) EXPECT() *MockRepositoryMockRecorder { + return m.recorder +} + +// GetAccountByPeerID mocks base method. +func (m *MockRepository) GetAccountByPeerID(ctx context.Context, peerID string) (*types.Account, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetAccountByPeerID", ctx, peerID) + ret0, _ := ret[0].(*types.Account) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetAccountByPeerID indicates an expected call of GetAccountByPeerID. +func (mr *MockRepositoryMockRecorder) GetAccountByPeerID(ctx, peerID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountByPeerID", reflect.TypeOf((*MockRepository)(nil).GetAccountByPeerID), ctx, peerID) +} + +// GetAccountNetwork mocks base method. +func (m *MockRepository) GetAccountNetwork(ctx context.Context, accountID string) (*types.Network, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetAccountNetwork", ctx, accountID) + ret0, _ := ret[0].(*types.Network) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetAccountNetwork indicates an expected call of GetAccountNetwork. +func (mr *MockRepositoryMockRecorder) GetAccountNetwork(ctx, accountID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountNetwork", reflect.TypeOf((*MockRepository)(nil).GetAccountNetwork), ctx, accountID) +} + +// GetAccountPeers mocks base method. +func (m *MockRepository) GetAccountPeers(ctx context.Context, accountID string) ([]*peer.Peer, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetAccountPeers", ctx, accountID) + ret0, _ := ret[0].([]*peer.Peer) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetAccountPeers indicates an expected call of GetAccountPeers. +func (mr *MockRepositoryMockRecorder) GetAccountPeers(ctx, accountID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountPeers", reflect.TypeOf((*MockRepository)(nil).GetAccountPeers), ctx, accountID) +} + +// GetAccountZones mocks base method. +func (m *MockRepository) GetAccountZones(ctx context.Context, accountID string) ([]*zones.Zone, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetAccountZones", ctx, accountID) + ret0, _ := ret[0].([]*zones.Zone) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetAccountZones indicates an expected call of GetAccountZones. +func (mr *MockRepositoryMockRecorder) GetAccountZones(ctx, accountID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountZones", reflect.TypeOf((*MockRepository)(nil).GetAccountZones), ctx, accountID) +} + +// GetPeerByID mocks base method. +func (m *MockRepository) GetPeerByID(ctx context.Context, accountID, peerID string) (*peer.Peer, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetPeerByID", ctx, accountID, peerID) + ret0, _ := ret[0].(*peer.Peer) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetPeerByID indicates an expected call of GetPeerByID. +func (mr *MockRepositoryMockRecorder) GetPeerByID(ctx, accountID, peerID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeerByID", reflect.TypeOf((*MockRepository)(nil).GetPeerByID), ctx, accountID, peerID) +} + +// GetPeersByIDs mocks base method. +func (m *MockRepository) GetPeersByIDs(ctx context.Context, accountID string, peerIDs []string) (map[string]*peer.Peer, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetPeersByIDs", ctx, accountID, peerIDs) + ret0, _ := ret[0].(map[string]*peer.Peer) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetPeersByIDs indicates an expected call of GetPeersByIDs. +func (mr *MockRepositoryMockRecorder) GetPeersByIDs(ctx, accountID, peerIDs any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeersByIDs", reflect.TypeOf((*MockRepository)(nil).GetPeersByIDs), ctx, accountID, peerIDs) +} + +// SynthesizeAgentNetworkServices mocks base method. +func (m *MockRepository) SynthesizeAgentNetworkServices(ctx context.Context, accountID string) ([]*service.Service, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "SynthesizeAgentNetworkServices", ctx, accountID) + ret0, _ := ret[0].([]*service.Service) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// SynthesizeAgentNetworkServices indicates an expected call of SynthesizeAgentNetworkServices. +func (mr *MockRepositoryMockRecorder) SynthesizeAgentNetworkServices(ctx, accountID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SynthesizeAgentNetworkServices", reflect.TypeOf((*MockRepository)(nil).SynthesizeAgentNetworkServices), ctx, accountID) +} diff --git a/management/server/account/request_buffer.go b/management/server/account/request_buffer.go index eced1929f..3f91996eb 100644 --- a/management/server/account/request_buffer.go +++ b/management/server/account/request_buffer.go @@ -6,6 +6,8 @@ import ( "github.com/netbirdio/netbird/management/server/types" ) +//go:generate go tool mockgen -package=account -source=./request_buffer.go -destination=request_buffer_mock.go + type RequestBuffer interface { GetAccountWithBackpressure(ctx context.Context, accountID string) (*types.Account, error) } diff --git a/management/server/account/request_buffer_mock.go b/management/server/account/request_buffer_mock.go new file mode 100644 index 000000000..b48ef2700 --- /dev/null +++ b/management/server/account/request_buffer_mock.go @@ -0,0 +1,57 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: ./request_buffer.go +// +// Generated by this command: +// +// mockgen -package=account -source=./request_buffer.go -destination=request_buffer_mock.go +// + +// Package account is a generated GoMock package. +package account + +import ( + context "context" + reflect "reflect" + + types "github.com/netbirdio/netbird/management/server/types" + gomock "go.uber.org/mock/gomock" +) + +// MockRequestBuffer is a mock of RequestBuffer interface. +type MockRequestBuffer struct { + ctrl *gomock.Controller + recorder *MockRequestBufferMockRecorder + isgomock struct{} +} + +// MockRequestBufferMockRecorder is the mock recorder for MockRequestBuffer. +type MockRequestBufferMockRecorder struct { + mock *MockRequestBuffer +} + +// NewMockRequestBuffer creates a new mock instance. +func NewMockRequestBuffer(ctrl *gomock.Controller) *MockRequestBuffer { + mock := &MockRequestBuffer{ctrl: ctrl} + mock.recorder = &MockRequestBufferMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockRequestBuffer) EXPECT() *MockRequestBufferMockRecorder { + return m.recorder +} + +// GetAccountWithBackpressure mocks base method. +func (m *MockRequestBuffer) GetAccountWithBackpressure(ctx context.Context, accountID string) (*types.Account, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetAccountWithBackpressure", ctx, accountID) + ret0, _ := ret[0].(*types.Account) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetAccountWithBackpressure indicates an expected call of GetAccountWithBackpressure. +func (mr *MockRequestBufferMockRecorder) GetAccountWithBackpressure(ctx, accountID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountWithBackpressure", reflect.TypeOf((*MockRequestBuffer)(nil).GetAccountWithBackpressure), ctx, accountID) +} diff --git a/management/server/types/account_components.go b/management/server/types/account_components.go index 624a778fe..3f2d5485f 100644 --- a/management/server/types/account_components.go +++ b/management/server/types/account_components.go @@ -112,8 +112,6 @@ func (a *Account) GetPeerNetworkMapComponents( return EmptyNetworkMapComponents(&NetworkMapComponents{ PeerID: peerID, Network: a.Network.Copy(), - // must include the target peer as it's required on the client - Peers: map[string]*ComponentPeer{peerID: peer.ToComponent()}, }) } diff --git a/management/server/types/account_components_test.go b/management/server/types/account_components_test.go new file mode 100644 index 000000000..3574480e8 --- /dev/null +++ b/management/server/types/account_components_test.go @@ -0,0 +1,20 @@ +package types + +import ( + "context" + "testing" + + "github.com/netbirdio/netbird/dns" + "github.com/netbirdio/netbird/shared/management/types" + "github.com/stretchr/testify/assert" +) + +func TestGetPeerNetworkMapComponents_PeerMissingFromAcount(t *testing.T) { + account := Account{Network: NewNetwork()} + nmapcomponets := account.GetPeerNetworkMapComponents(context.TODO(), "missing-peer", dns.CustomZone{}, nil, nil, nil, nil, nil) + + assert.Equal(t, EmptyNetworkMapComponents(&types.NetworkMapComponents{ + PeerID: "missing-peer", + Network: account.Network, + }), nmapcomponets) +}