From 285bbc5ffb64860ba334dd39dba405340fea35b7 Mon Sep 17 00:00:00 2001 From: pascal Date: Mon, 27 Apr 2026 17:49:12 +0200 Subject: [PATCH 01/28] calculate affected peers --- .../network_map/controller/controller.go | 131 ++++++++ .../controllers/network_map/interface.go | 1 + .../controllers/network_map/interface_mock.go | 14 + management/server/account/manager.go | 1 + management/server/account/manager_mock.go | 12 + management/server/dns.go | 28 +- management/server/group.go | 282 +++++++++++++----- management/server/mock_server/account_mock.go | 7 + management/server/nameserver.go | 57 +--- management/server/peer.go | 32 ++ management/server/policy.go | 86 ++---- management/server/posture_checks.go | 34 +-- management/server/route.go | 80 ++--- management/server/store/sql_store.go | 17 ++ management/server/store/store.go | 1 + management/server/store/store_mock.go | 16 + 16 files changed, 540 insertions(+), 259 deletions(-) diff --git a/management/internals/controllers/network_map/controller/controller.go b/management/internals/controllers/network_map/controller/controller.go index 4b47ecaa0..5cbf6fceb 100644 --- a/management/internals/controllers/network_map/controller/controller.go +++ b/management/internals/controllers/network_map/controller/controller.go @@ -261,6 +261,137 @@ func (c *Controller) UpdateAccountPeers(ctx context.Context, accountID string) e return c.sendUpdateAccountPeers(ctx, accountID) } +// UpdateAffectedPeers updates only the specified peers that belong to an account. +// Should be called when a change is known to affect only a subset of peers. +// If peerIDs is empty, this is a no-op. +func (c *Controller) UpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error { + if len(peerIDs) == 0 { + return nil + } + return c.sendUpdateForAffectedPeers(ctx, accountID, peerIDs) +} + +func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error { + log.WithContext(ctx).Tracef("updating %d affected peers for account %s from %s", len(peerIDs), accountID, util.GetCallerName()) + + affected := make(map[string]struct{}, len(peerIDs)) + for _, id := range peerIDs { + affected[id] = struct{}{} + } + + // Fast check: any of the affected peers actually connected? + hasConnected := false + for _, id := range peerIDs { + if c.peersUpdateManager.HasChannel(id) { + hasConnected = true + break + } + } + if !hasConnected { + return nil + } + + account, err := c.requestBuffer.GetAccountWithBackpressure(ctx, accountID) + if err != nil { + return fmt.Errorf("failed to get account: %v", err) + } + + globalStart := time.Now() + + // Collect the subset of account peers that are both affected and connected. + var peersToUpdate []*nbpeer.Peer + for _, peer := range account.Peers { + if _, ok := affected[peer.ID]; ok && c.peersUpdateManager.HasChannel(peer.ID) { + peersToUpdate = append(peersToUpdate, peer) + } + } + + if len(peersToUpdate) == 0 { + return nil + } + + approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra) + if err != nil { + return fmt.Errorf("failed to get validate peers: %v", err) + } + + var wg sync.WaitGroup + semaphore := make(chan struct{}, 10) + + account.InjectProxyPolicies(ctx) + dnsCache := &cache.DNSConfigCache{} + dnsDomain := c.GetDNSDomain(account.Settings) + peersCustomZone := account.GetPeersCustomZone(ctx, dnsDomain) + resourcePolicies := account.GetResourcePoliciesMap() + routers := account.GetResourceRoutersMap() + groupIDToUserIDs := account.GetActiveGroupUsers() + + proxyNetworkMaps, err := c.proxyController.GetProxyNetworkMapsAll(ctx, accountID, account.Peers) + if err != nil { + log.WithContext(ctx).Errorf("failed to get proxy network maps: %v", err) + return fmt.Errorf("failed to get proxy network maps: %v", err) + } + + extraSetting, err := c.settingsManager.GetExtraSettings(ctx, accountID) + if err != nil { + return fmt.Errorf("failed to get flow enabled status: %v", err) + } + + dnsFwdPort := computeForwarderPort(maps.Values(account.Peers), network_map.DnsForwarderPortMinVersion) + + accountZones, err := c.repo.GetAccountZones(ctx, account.Id) + if err != nil { + log.WithContext(ctx).Errorf("failed to get account zones: %v", err) + return fmt.Errorf("failed to get account zones: %v", err) + } + + for _, peer := range peersToUpdate { + wg.Add(1) + semaphore <- struct{}{} + go func(p *nbpeer.Peer) { + defer wg.Done() + defer func() { <-semaphore }() + + start := time.Now() + + postureChecks, err := c.getPeerPostureChecks(account, p.ID) + if err != nil { + log.WithContext(ctx).Debugf("failed to get posture checks for peer %s: %v", p.ID, err) + return + } + + c.metrics.CountCalcPostureChecksDuration(time.Since(start)) + start = time.Now() + + remotePeerNetworkMap := account.GetPeerNetworkMapFromComponents(ctx, p.ID, peersCustomZone, accountZones, approvedPeersMap, resourcePolicies, routers, c.accountManagerMetrics, groupIDToUserIDs) + + c.metrics.CountCalcPeerNetworkMapDuration(time.Since(start)) + + proxyNetworkMap, ok := proxyNetworkMaps[p.ID] + if ok { + remotePeerNetworkMap.Merge(proxyNetworkMap) + } + + peerGroups := account.GetPeerGroups(p.ID) + start = time.Now() + update := grpc.ToSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, p, nil, nil, remotePeerNetworkMap, dnsDomain, postureChecks, dnsCache, account.Settings, extraSetting, maps.Keys(peerGroups), dnsFwdPort) + c.metrics.CountToSyncResponseDuration(time.Since(start)) + + c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{ + Update: update, + MessageType: network_map.MessageTypeNetworkMap, + }) + }(peer) + } + + wg.Wait() + if c.accountManagerMetrics != nil { + c.accountManagerMetrics.CountUpdateAccountPeersDuration(time.Since(globalStart)) + } + + return nil +} + func (c *Controller) UpdateAccountPeer(ctx context.Context, accountId string, peerId string) error { if !c.peersUpdateManager.HasChannel(peerId) { return fmt.Errorf("peer %s doesn't have a channel, skipping network map update", peerId) diff --git a/management/internals/controllers/network_map/interface.go b/management/internals/controllers/network_map/interface.go index cfea2d3de..8d81556f9 100644 --- a/management/internals/controllers/network_map/interface.go +++ b/management/internals/controllers/network_map/interface.go @@ -19,6 +19,7 @@ const ( type Controller interface { UpdateAccountPeers(ctx context.Context, accountID string) error + UpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error UpdateAccountPeer(ctx context.Context, accountId string, peerId string) error BufferUpdateAccountPeers(ctx context.Context, accountID string) error GetValidatedPeerWithMap(ctx context.Context, isRequiresApproval bool, accountID string, p *nbpeer.Peer) (*nbpeer.Peer, *types.NetworkMap, []*posture.Checks, int64, error) diff --git a/management/internals/controllers/network_map/interface_mock.go b/management/internals/controllers/network_map/interface_mock.go index 4e86d2973..b2ef0b861 100644 --- a/management/internals/controllers/network_map/interface_mock.go +++ b/management/internals/controllers/network_map/interface_mock.go @@ -250,3 +250,17 @@ func (mr *MockControllerMockRecorder) UpdateAccountPeers(ctx, accountID any) *go mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateAccountPeers", reflect.TypeOf((*MockController)(nil).UpdateAccountPeers), ctx, accountID) } + +// UpdateAffectedPeers mocks base method. +func (m *MockController) UpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "UpdateAffectedPeers", ctx, accountID, peerIDs) + ret0, _ := ret[0].(error) + return ret0 +} + +// UpdateAffectedPeers indicates an expected call of UpdateAffectedPeers. +func (mr *MockControllerMockRecorder) UpdateAffectedPeers(ctx, accountID, peerIDs any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateAffectedPeers", reflect.TypeOf((*MockController)(nil).UpdateAffectedPeers), ctx, accountID, peerIDs) +} diff --git a/management/server/account/manager.go b/management/server/account/manager.go index b4516d512..576054d1e 100644 --- a/management/server/account/manager.go +++ b/management/server/account/manager.go @@ -125,6 +125,7 @@ type Manager interface { GetAccountSettings(ctx context.Context, accountID string, userID string) (*types.Settings, error) DeleteSetupKey(ctx context.Context, accountID, userID, keyID string) error UpdateAccountPeers(ctx context.Context, accountID string) + UpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) BufferUpdateAccountPeers(ctx context.Context, accountID string) BuildUserInfosForAccount(ctx context.Context, accountID, initiatorUserID string, accountUsers []*types.User) (map[string]*types.UserInfo, error) SyncUserJWTGroups(ctx context.Context, userAuth auth.UserAuth) error diff --git a/management/server/account/manager_mock.go b/management/server/account/manager_mock.go index 36e5fe39f..c595346f8 100644 --- a/management/server/account/manager_mock.go +++ b/management/server/account/manager_mock.go @@ -1608,6 +1608,18 @@ func (mr *MockManagerMockRecorder) UpdateAccountPeers(ctx, accountID interface{} return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateAccountPeers", reflect.TypeOf((*MockManager)(nil).UpdateAccountPeers), ctx, accountID) } +// UpdateAffectedPeers mocks base method. +func (m *MockManager) UpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "UpdateAffectedPeers", ctx, accountID, peerIDs) +} + +// UpdateAffectedPeers indicates an expected call of UpdateAffectedPeers. +func (mr *MockManagerMockRecorder) UpdateAffectedPeers(ctx, accountID, peerIDs interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateAffectedPeers", reflect.TypeOf((*MockManager)(nil).UpdateAffectedPeers), ctx, accountID, peerIDs) +} + // UpdateAccountSettings mocks base method. func (m *MockManager) UpdateAccountSettings(ctx context.Context, accountID, userID string, newSettings *types.Settings) (*types.Settings, error) { m.ctrl.T.Helper() diff --git a/management/server/dns.go b/management/server/dns.go index baf6debc3..1e213ffbb 100644 --- a/management/server/dns.go +++ b/management/server/dns.go @@ -47,8 +47,8 @@ func (am *DefaultAccountManager) SaveDNSSettings(ctx context.Context, accountID return status.NewPermissionDeniedError() } - var updateAccountPeers bool var eventsToStore []func() + var affectedPeerIDs []string err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { if err = validateDNSSettings(ctx, transaction, accountID, dnsSettingsToSave); err != nil { @@ -63,11 +63,6 @@ func (am *DefaultAccountManager) SaveDNSSettings(ctx context.Context, accountID addedGroups := util.Difference(dnsSettingsToSave.DisabledManagementGroups, oldSettings.DisabledManagementGroups) removedGroups := util.Difference(oldSettings.DisabledManagementGroups, dnsSettingsToSave.DisabledManagementGroups) - updateAccountPeers, err = areDNSSettingChangesAffectPeers(ctx, transaction, accountID, addedGroups, removedGroups) - if err != nil { - return err - } - events := am.prepareDNSSettingsEvents(ctx, transaction, accountID, userID, addedGroups, removedGroups) eventsToStore = append(eventsToStore, events...) @@ -75,6 +70,9 @@ func (am *DefaultAccountManager) SaveDNSSettings(ctx context.Context, accountID return err } + allGroups := slices.Concat(addedGroups, removedGroups) + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, allGroups, nil) + return transaction.IncrementNetworkSerial(ctx, accountID) }) if err != nil { @@ -85,8 +83,8 @@ func (am *DefaultAccountManager) SaveDNSSettings(ctx context.Context, accountID storeEvent() } - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return nil @@ -133,20 +131,6 @@ func (am *DefaultAccountManager) prepareDNSSettingsEvents(ctx context.Context, t return eventsToStore } -// areDNSSettingChangesAffectPeers checks if the DNS settings changes affect any peers. -func areDNSSettingChangesAffectPeers(ctx context.Context, transaction store.Store, accountID string, addedGroups, removedGroups []string) (bool, error) { - hasPeers, err := anyGroupHasPeersOrResources(ctx, transaction, accountID, addedGroups) - if err != nil { - return false, err - } - - if hasPeers { - return true, nil - } - - return anyGroupHasPeersOrResources(ctx, transaction, accountID, removedGroups) -} - // validateDNSSettings validates the DNS settings. func validateDNSSettings(ctx context.Context, transaction store.Store, accountID string, settings *types.DNSSettings) error { if len(settings.DisabledManagementGroups) == 0 { diff --git a/management/server/group.go b/management/server/group.go index 7b5b9b86c..4bd249398 100644 --- a/management/server/group.go +++ b/management/server/group.go @@ -79,7 +79,7 @@ func (am *DefaultAccountManager) CreateGroup(ctx context.Context, accountID, use } var eventsToStore []func() - var updateAccountPeers bool + var affectedPeerIDs []string err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { if err = validateNewGroup(ctx, transaction, accountID, newGroup); err != nil { @@ -91,11 +91,6 @@ func (am *DefaultAccountManager) CreateGroup(ctx context.Context, accountID, use events := am.prepareGroupEvents(ctx, transaction, accountID, userID, newGroup) eventsToStore = append(eventsToStore, events...) - updateAccountPeers, err = areGroupChangesAffectPeers(ctx, transaction, accountID, []string{newGroup.ID}) - if err != nil { - return err - } - if err := transaction.CreateGroup(ctx, newGroup); err != nil { return status.Errorf(status.Internal, "failed to create group: %v", err) } @@ -106,6 +101,9 @@ func (am *DefaultAccountManager) CreateGroup(ctx context.Context, accountID, use } } + groupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, transaction, accountID, []string{newGroup.ID}) + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) + return transaction.IncrementNetworkSerial(ctx, accountID) }) if err != nil { @@ -116,8 +114,8 @@ func (am *DefaultAccountManager) CreateGroup(ctx context.Context, accountID, use storeEvent() } - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return nil @@ -134,7 +132,7 @@ func (am *DefaultAccountManager) UpdateGroup(ctx context.Context, accountID, use } var eventsToStore []func() - var updateAccountPeers bool + var affectedPeerIDs []string err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { if err = validateNewGroup(ctx, transaction, accountID, newGroup); err != nil { @@ -165,15 +163,13 @@ func (am *DefaultAccountManager) UpdateGroup(ctx context.Context, accountID, use } } - updateAccountPeers, err = areGroupChangesAffectPeers(ctx, transaction, accountID, []string{newGroup.ID}) - if err != nil { - return err - } - if err = transaction.UpdateGroup(ctx, newGroup); err != nil { return err } + groupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, transaction, accountID, []string{newGroup.ID}) + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) + return transaction.IncrementNetworkSerial(ctx, accountID) }) if err != nil { @@ -184,8 +180,8 @@ func (am *DefaultAccountManager) UpdateGroup(ctx context.Context, accountID, use storeEvent() } - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return nil @@ -205,7 +201,6 @@ func (am *DefaultAccountManager) CreateGroups(ctx context.Context, accountID, us } var eventsToStore []func() - var updateAccountPeers bool var globalErr error groupIDs := make([]string, 0, len(groups)) @@ -243,17 +238,14 @@ func (am *DefaultAccountManager) CreateGroups(ctx context.Context, accountID, us } } - updateAccountPeers, err = areGroupChangesAffectPeers(ctx, am.Store, accountID, groupIDs) - if err != nil { - return err - } - for _, storeEvent := range eventsToStore { storeEvent() } - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, am.Store, accountID, groupIDs) + affectedPeerIDs := am.resolvePeerIDs(ctx, am.Store, accountID, allGroupIDs, directPeerIDs) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return globalErr @@ -273,7 +265,6 @@ func (am *DefaultAccountManager) UpdateGroups(ctx context.Context, accountID, us } var eventsToStore []func() - var updateAccountPeers bool var globalErr error groupIDs := make([]string, 0, len(groups)) @@ -311,17 +302,14 @@ func (am *DefaultAccountManager) UpdateGroups(ctx context.Context, accountID, us } } - updateAccountPeers, err = areGroupChangesAffectPeers(ctx, am.Store, accountID, groupIDs) - if err != nil { - return err - } - for _, storeEvent := range eventsToStore { storeEvent() } - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, am.Store, accountID, groupIDs) + affectedPeerIDs := am.resolvePeerIDs(ctx, am.Store, accountID, allGroupIDs, directPeerIDs) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return globalErr @@ -473,27 +461,25 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us // GroupAddPeer appends peer to the group func (am *DefaultAccountManager) GroupAddPeer(ctx context.Context, accountID, groupID, peerID string) error { - var updateAccountPeers bool + var affectedPeerIDs []string var err error err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { - updateAccountPeers, err = areGroupChangesAffectPeers(ctx, transaction, accountID, []string{groupID}) - if err != nil { - return err - } - if err = transaction.AddPeerToGroup(ctx, accountID, peerID, groupID); err != nil { return err } + allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, transaction, accountID, []string{groupID}) + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, allGroupIDs, directPeerIDs) + return transaction.IncrementNetworkSerial(ctx, accountID) }) if err != nil { return err } - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return nil @@ -502,7 +488,7 @@ func (am *DefaultAccountManager) GroupAddPeer(ctx context.Context, accountID, gr // GroupAddResource appends resource to the group func (am *DefaultAccountManager) GroupAddResource(ctx context.Context, accountID, groupID string, resource types.Resource) error { var group *types.Group - var updateAccountPeers bool + var affectedPeerIDs []string var err error err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { @@ -515,23 +501,21 @@ func (am *DefaultAccountManager) GroupAddResource(ctx context.Context, accountID return nil } - updateAccountPeers, err = areGroupChangesAffectPeers(ctx, transaction, accountID, []string{groupID}) - if err != nil { - return err - } - if err = transaction.UpdateGroup(ctx, group); err != nil { return err } + allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, transaction, accountID, []string{groupID}) + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, allGroupIDs, directPeerIDs) + return transaction.IncrementNetworkSerial(ctx, accountID) }) if err != nil { return err } - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return nil @@ -539,14 +523,13 @@ func (am *DefaultAccountManager) GroupAddResource(ctx context.Context, accountID // GroupDeletePeer removes peer from the group func (am *DefaultAccountManager) GroupDeletePeer(ctx context.Context, accountID, groupID, peerID string) error { - var updateAccountPeers bool + var affectedPeerIDs []string var err error err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { - updateAccountPeers, err = areGroupChangesAffectPeers(ctx, transaction, accountID, []string{groupID}) - if err != nil { - return err - } + // Resolve before removing, so the peer being removed is still included + allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, transaction, accountID, []string{groupID}) + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, allGroupIDs, directPeerIDs) if err = transaction.RemovePeerFromGroup(ctx, peerID, groupID); err != nil { return err @@ -558,8 +541,8 @@ func (am *DefaultAccountManager) GroupDeletePeer(ctx context.Context, accountID, return err } - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return nil @@ -568,7 +551,7 @@ func (am *DefaultAccountManager) GroupDeletePeer(ctx context.Context, accountID, // GroupDeleteResource removes resource from the group func (am *DefaultAccountManager) GroupDeleteResource(ctx context.Context, accountID, groupID string, resource types.Resource) error { var group *types.Group - var updateAccountPeers bool + var affectedPeerIDs []string var err error err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { @@ -581,23 +564,21 @@ func (am *DefaultAccountManager) GroupDeleteResource(ctx context.Context, accoun return nil } - updateAccountPeers, err = areGroupChangesAffectPeers(ctx, transaction, accountID, []string{groupID}) - if err != nil { - return err - } - if err = transaction.UpdateGroup(ctx, group); err != nil { return err } + allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, transaction, accountID, []string{groupID}) + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, allGroupIDs, directPeerIDs) + return transaction.IncrementNetworkSerial(ctx, accountID) }) if err != nil { return err } - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return nil @@ -840,18 +821,175 @@ func areGroupChangesAffectPeers(ctx context.Context, transaction store.Store, ac return false, nil } -// anyGroupHasPeersOrResources checks if any of the given groups in the account have peers or resources. -func anyGroupHasPeersOrResources(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) (bool, error) { - groups, err := transaction.GetGroupsByIDs(ctx, store.LockingStrengthNone, accountID, groupIDs) - if err != nil { - return false, err +// collectGroupChangeAffectedGroups walks all entities that reference the changed groups +// and collects the full set of affected group IDs and direct peer IDs. +// This ensures that when a group changes, we update not just the peers in that group +// but also peers in other groups that share policies, routes, DNS, or nameserver configs. +func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedGroupIDs []string) (allGroupIDs []string, directPeerIDs []string) { + if len(changedGroupIDs) == 0 { + return nil, nil } - for _, group := range groups { - if group.HasPeers() || group.HasResources() { - return true, nil + changedSet := make(map[string]struct{}, len(changedGroupIDs)) + for _, id := range changedGroupIDs { + changedSet[id] = struct{}{} + } + + groupSet := make(map[string]struct{}) + // Always include the changed groups themselves + for _, id := range changedGroupIDs { + groupSet[id] = struct{}{} + } + + peerSet := make(map[string]struct{}) + + // Policies: collect all rule groups + direct peer resources from policies that reference any changed group + policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get policies for group change resolution: %v", err) + } else { + for _, policy := range policies { + if !policyReferencesGroups(policy, changedSet) { + continue + } + for _, gID := range policy.RuleGroups() { + groupSet[gID] = struct{}{} + } + for _, rule := range policy.Rules { + if rule.SourceResource.Type == types.ResourceTypePeer && rule.SourceResource.ID != "" { + peerSet[rule.SourceResource.ID] = struct{}{} + } + if rule.DestinationResource.Type == types.ResourceTypePeer && rule.DestinationResource.ID != "" { + peerSet[rule.DestinationResource.ID] = struct{}{} + } + } } } - return false, nil + // Routes: collect all groups + direct peer from routes that reference any changed group + routes, err := transaction.GetAccountRoutes(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get routes for group change resolution: %v", err) + } else { + for _, r := range routes { + if !routeReferencesGroups(r, changedSet) { + continue + } + for _, gID := range r.Groups { + groupSet[gID] = struct{}{} + } + for _, gID := range r.PeerGroups { + groupSet[gID] = struct{}{} + } + for _, gID := range r.AccessControlGroups { + groupSet[gID] = struct{}{} + } + if r.Peer != "" { + peerSet[r.Peer] = struct{}{} + } + } + } + + // Nameserver groups: collect groups from NS groups that reference any changed group + nsGroups, err := transaction.GetAccountNameServerGroups(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get nameserver groups for group change resolution: %v", err) + } else { + for _, ns := range nsGroups { + for _, gID := range ns.Groups { + if _, ok := changedSet[gID]; ok { + for _, g := range ns.Groups { + groupSet[g] = struct{}{} + } + break + } + } + } + } + + // DNS settings: if any changed group is in DisabledManagementGroups, include those groups + dnsSettings, err := transaction.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get DNS settings for group change resolution: %v", err) + } else { + for _, gID := range dnsSettings.DisabledManagementGroups { + if _, ok := changedSet[gID]; ok { + groupSet[gID] = struct{}{} + } + } + } + + // Network routers: collect peer groups + direct peer from routers that reference any changed group + routers, err := transaction.GetNetworkRoutersByAccountID(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get network routers for group change resolution: %v", err) + } else { + for _, router := range routers { + if !routerReferencesGroups(router, changedSet) { + continue + } + for _, gID := range router.PeerGroups { + groupSet[gID] = struct{}{} + } + if router.Peer != "" { + peerSet[router.Peer] = struct{}{} + } + } + } + + allGroupIDs = make([]string, 0, len(groupSet)) + for gID := range groupSet { + allGroupIDs = append(allGroupIDs, gID) + } + + directPeerIDs = make([]string, 0, len(peerSet)) + for pID := range peerSet { + directPeerIDs = append(directPeerIDs, pID) + } + + return allGroupIDs, directPeerIDs +} + +func policyReferencesGroups(policy *types.Policy, groupSet map[string]struct{}) bool { + for _, rule := range policy.Rules { + for _, gID := range rule.Sources { + if _, ok := groupSet[gID]; ok { + return true + } + } + for _, gID := range rule.Destinations { + if _, ok := groupSet[gID]; ok { + return true + } + } + } + return false +} + +func routeReferencesGroups(r *route.Route, groupSet map[string]struct{}) bool { + for _, gID := range r.Groups { + if _, ok := groupSet[gID]; ok { + return true + } + } + for _, gID := range r.PeerGroups { + if _, ok := groupSet[gID]; ok { + return true + } + } + for _, gID := range r.AccessControlGroups { + if _, ok := groupSet[gID]; ok { + return true + } + } + return false +} + +func routerReferencesGroups(router *routerTypes.NetworkRouter, groupSet map[string]struct{}) bool { + for _, gID := range router.PeerGroups { + if _, ok := groupSet[gID]; ok { + return true + } + } + return false } diff --git a/management/server/mock_server/account_mock.go b/management/server/mock_server/account_mock.go index ff369355e..5a9009da7 100644 --- a/management/server/mock_server/account_mock.go +++ b/management/server/mock_server/account_mock.go @@ -129,6 +129,7 @@ type MockAccountManager struct { AllowSyncFunc func(string, uint64) bool UpdateAccountPeersFunc func(ctx context.Context, accountID string) + UpdateAffectedPeersFunc func(ctx context.Context, accountID string, peerIDs []string) BufferUpdateAccountPeersFunc func(ctx context.Context, accountID string) RecalculateNetworkMapCacheFunc func(ctx context.Context, accountId string) error @@ -206,6 +207,12 @@ func (am *MockAccountManager) UpdateAccountPeers(ctx context.Context, accountID } } +func (am *MockAccountManager) UpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) { + if am.UpdateAffectedPeersFunc != nil { + am.UpdateAffectedPeersFunc(ctx, accountID, peerIDs) + } +} + func (am *MockAccountManager) BufferUpdateAccountPeers(ctx context.Context, accountID string) { if am.BufferUpdateAccountPeersFunc != nil { am.BufferUpdateAccountPeersFunc(ctx, accountID) diff --git a/management/server/nameserver.go b/management/server/nameserver.go index 3d8c78912..823fc72d5 100644 --- a/management/server/nameserver.go +++ b/management/server/nameserver.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "slices" "strings" "unicode/utf8" @@ -57,22 +58,19 @@ func (am *DefaultAccountManager) CreateNameServerGroup(ctx context.Context, acco SearchDomainsEnabled: searchDomainEnabled, } - var updateAccountPeers bool + var affectedPeerIDs []string err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { if err = validateNameServerGroup(ctx, transaction, accountID, newNSGroup); err != nil { return err } - updateAccountPeers, err = anyGroupHasPeersOrResources(ctx, transaction, accountID, newNSGroup.Groups) - if err != nil { - return err - } - if err = transaction.SaveNameServerGroup(ctx, newNSGroup); err != nil { return err } + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, newNSGroup.Groups, nil) + return transaction.IncrementNetworkSerial(ctx, accountID) }) if err != nil { @@ -81,8 +79,8 @@ func (am *DefaultAccountManager) CreateNameServerGroup(ctx context.Context, acco am.StoreEvent(ctx, userID, newNSGroup.ID, accountID, activity.NameserverGroupCreated, newNSGroup.EventMeta()) - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return newNSGroup.Copy(), nil @@ -102,7 +100,7 @@ func (am *DefaultAccountManager) SaveNameServerGroup(ctx context.Context, accoun return status.NewPermissionDeniedError() } - var updateAccountPeers bool + var affectedPeerIDs []string err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { oldNSGroup, err := transaction.GetNameServerGroupByID(ctx, store.LockingStrengthNone, accountID, nsGroupToSave.ID) @@ -115,15 +113,13 @@ func (am *DefaultAccountManager) SaveNameServerGroup(ctx context.Context, accoun return err } - updateAccountPeers, err = areNameServerGroupChangesAffectPeers(ctx, transaction, nsGroupToSave, oldNSGroup) - if err != nil { - return err - } - if err = transaction.SaveNameServerGroup(ctx, nsGroupToSave); err != nil { return err } + allGroups := slices.Concat(nsGroupToSave.Groups, oldNSGroup.Groups) + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, allGroups, nil) + return transaction.IncrementNetworkSerial(ctx, accountID) }) if err != nil { @@ -132,8 +128,8 @@ func (am *DefaultAccountManager) SaveNameServerGroup(ctx context.Context, accoun am.StoreEvent(ctx, userID, nsGroupToSave.ID, accountID, activity.NameserverGroupUpdated, nsGroupToSave.EventMeta()) - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return nil @@ -150,7 +146,7 @@ func (am *DefaultAccountManager) DeleteNameServerGroup(ctx context.Context, acco } var nsGroup *nbdns.NameServerGroup - var updateAccountPeers bool + var affectedPeerIDs []string err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { nsGroup, err = transaction.GetNameServerGroupByID(ctx, store.LockingStrengthUpdate, accountID, nsGroupID) @@ -158,10 +154,7 @@ func (am *DefaultAccountManager) DeleteNameServerGroup(ctx context.Context, acco return err } - updateAccountPeers, err = anyGroupHasPeersOrResources(ctx, transaction, accountID, nsGroup.Groups) - if err != nil { - return err - } + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, nsGroup.Groups, nil) if err = transaction.DeleteNameServerGroup(ctx, accountID, nsGroupID); err != nil { return err @@ -175,8 +168,8 @@ func (am *DefaultAccountManager) DeleteNameServerGroup(ctx context.Context, acco am.StoreEvent(ctx, userID, nsGroup.ID, accountID, activity.NameserverGroupDeleted, nsGroup.EventMeta()) - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return nil @@ -224,24 +217,6 @@ func validateNameServerGroup(ctx context.Context, transaction store.Store, accou return validateGroups(nameserverGroup.Groups, groups) } -// areNameServerGroupChangesAffectPeers checks if the changes in the nameserver group affect the peers. -func areNameServerGroupChangesAffectPeers(ctx context.Context, transaction store.Store, newNSGroup, oldNSGroup *nbdns.NameServerGroup) (bool, error) { - if !newNSGroup.Enabled && !oldNSGroup.Enabled { - return false, nil - } - - hasPeers, err := anyGroupHasPeersOrResources(ctx, transaction, newNSGroup.AccountID, newNSGroup.Groups) - if err != nil { - return false, err - } - - if hasPeers { - return true, nil - } - - return anyGroupHasPeersOrResources(ctx, transaction, oldNSGroup.AccountID, oldNSGroup.Groups) -} - func validateDomainInput(primary bool, domains []string, searchDomainsEnabled bool) error { if !primary && len(domains) == 0 { return status.Errorf(status.InvalidArgument, "nameserver group primary status is false and domains are empty,"+ diff --git a/management/server/peer.go b/management/server/peer.go index a95ae17a3..39368c840 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -1294,6 +1294,38 @@ func (am *DefaultAccountManager) UpdateAccountPeers(ctx context.Context, account _ = am.networkMapController.UpdateAccountPeers(ctx, accountID) } +// UpdateAffectedPeers updates only the specified peers that belong to an account. +// Should be called when a change is known to affect only a subset of peers. +func (am *DefaultAccountManager) UpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) { + _ = am.networkMapController.UpdateAffectedPeers(ctx, accountID, peerIDs) +} + +// resolvePeerIDs resolves a set of group IDs and direct peer IDs into a +// deduplicated list of peer IDs suitable for UpdateAffectedPeers. +func (am *DefaultAccountManager) resolvePeerIDs(ctx context.Context, s store.Store, accountID string, groupIDs []string, directPeerIDs []string) []string { + peerIDs, err := s.GetPeerIDsByGroups(ctx, accountID, groupIDs) + if err != nil { + log.WithContext(ctx).Errorf("failed to resolve peer IDs by groups: %v", err) + return nil + } + + if len(directPeerIDs) == 0 { + return peerIDs + } + + seen := make(map[string]struct{}, len(peerIDs)) + for _, id := range peerIDs { + seen[id] = struct{}{} + } + for _, id := range directPeerIDs { + if _, exists := seen[id]; !exists { + peerIDs = append(peerIDs, id) + seen[id] = struct{}{} + } + } + return peerIDs +} + func (am *DefaultAccountManager) BufferUpdateAccountPeers(ctx context.Context, accountID string) { _ = am.networkMapController.BufferUpdateAccountPeers(ctx, accountID) } diff --git a/management/server/policy.go b/management/server/policy.go index 48297ca11..9ba4f98dd 100644 --- a/management/server/policy.go +++ b/management/server/policy.go @@ -45,12 +45,13 @@ func (am *DefaultAccountManager) SavePolicy(ctx context.Context, accountID, user } var isUpdate = policy.ID != "" - var updateAccountPeers bool + var existingPolicy *types.Policy var action = activity.PolicyAdded var unchanged bool + var affectedPeerIDs []string err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { - existingPolicy, err := validatePolicy(ctx, transaction, accountID, policy) + existingPolicy, err = validatePolicy(ctx, transaction, accountID, policy) if err != nil { return err } @@ -64,25 +65,18 @@ func (am *DefaultAccountManager) SavePolicy(ctx context.Context, accountID, user action = activity.PolicyUpdated - updateAccountPeers, err = arePolicyChangesAffectPeersWithExisting(ctx, transaction, policy, existingPolicy) - if err != nil { - return err - } - if err = transaction.SavePolicy(ctx, policy); err != nil { return err } } else { - updateAccountPeers, err = arePolicyChangesAffectPeers(ctx, transaction, policy) - if err != nil { - return err - } - if err = transaction.CreatePolicy(ctx, policy); err != nil { return err } } + groupIDs, directPeerIDs := collectPolicyAffectedGroupsAndPeers(policy, existingPolicy) + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) + return transaction.IncrementNetworkSerial(ctx, accountID) }) if err != nil { @@ -95,8 +89,8 @@ func (am *DefaultAccountManager) SavePolicy(ctx context.Context, accountID, user am.StoreEvent(ctx, userID, policy.ID, accountID, action, policy.EventMeta()) - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return policy, nil @@ -113,7 +107,7 @@ func (am *DefaultAccountManager) DeletePolicy(ctx context.Context, accountID, po } var policy *types.Policy - var updateAccountPeers bool + var affectedPeerIDs []string err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { policy, err = transaction.GetPolicyByID(ctx, store.LockingStrengthUpdate, accountID, policyID) @@ -121,10 +115,8 @@ func (am *DefaultAccountManager) DeletePolicy(ctx context.Context, accountID, po return err } - updateAccountPeers, err = arePolicyChangesAffectPeers(ctx, transaction, policy) - if err != nil { - return err - } + groupIDs, directPeerIDs := collectPolicyAffectedGroupsAndPeers(policy) + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) if err = transaction.DeletePolicy(ctx, accountID, policyID); err != nil { return err @@ -138,8 +130,8 @@ func (am *DefaultAccountManager) DeletePolicy(ctx context.Context, accountID, po am.StoreEvent(ctx, userID, policyID, accountID, activity.PolicyRemoved, policy.EventMeta()) - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return nil @@ -158,44 +150,24 @@ func (am *DefaultAccountManager) ListPolicies(ctx context.Context, accountID, us return am.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) } -// arePolicyChangesAffectPeers checks if a policy (being created or deleted) will affect any associated peers. -func arePolicyChangesAffectPeers(ctx context.Context, transaction store.Store, policy *types.Policy) (bool, error) { - for _, rule := range policy.Rules { - if rule.SourceResource.Type != "" || rule.DestinationResource.Type != "" { - return true, nil +// collectPolicyAffectedGroupsAndPeers returns the group IDs and direct peer IDs +// referenced by the given policies' rules. +func collectPolicyAffectedGroupsAndPeers(policies ...*types.Policy) (groupIDs []string, directPeerIDs []string) { + for _, policy := range policies { + if policy == nil { + continue + } + groupIDs = append(groupIDs, policy.RuleGroups()...) + for _, rule := range policy.Rules { + if rule.SourceResource.Type == types.ResourceTypePeer && rule.SourceResource.ID != "" { + directPeerIDs = append(directPeerIDs, rule.SourceResource.ID) + } + if rule.DestinationResource.Type == types.ResourceTypePeer && rule.DestinationResource.ID != "" { + directPeerIDs = append(directPeerIDs, rule.DestinationResource.ID) + } } } - - return anyGroupHasPeersOrResources(ctx, transaction, policy.AccountID, policy.RuleGroups()) -} - -func arePolicyChangesAffectPeersWithExisting(ctx context.Context, transaction store.Store, policy *types.Policy, existingPolicy *types.Policy) (bool, error) { - if !policy.Enabled && !existingPolicy.Enabled { - return false, nil - } - - for _, rule := range existingPolicy.Rules { - if rule.SourceResource.Type != "" || rule.DestinationResource.Type != "" { - return true, nil - } - } - - hasPeers, err := anyGroupHasPeersOrResources(ctx, transaction, policy.AccountID, existingPolicy.RuleGroups()) - if err != nil { - return false, err - } - - if hasPeers { - return true, nil - } - - for _, rule := range policy.Rules { - if rule.SourceResource.Type != "" || rule.DestinationResource.Type != "" { - return true, nil - } - } - - return anyGroupHasPeersOrResources(ctx, transaction, policy.AccountID, policy.RuleGroups()) + return } // validatePolicy validates the policy and its rules. For updates it returns diff --git a/management/server/posture_checks.go b/management/server/posture_checks.go index 9562487c0..bbf4ed198 100644 --- a/management/server/posture_checks.go +++ b/management/server/posture_checks.go @@ -40,9 +40,9 @@ func (am *DefaultAccountManager) SavePostureChecks(ctx context.Context, accountI return nil, status.NewPermissionDeniedError() } - var updateAccountPeers bool var isUpdate = postureChecks.ID != "" var action = activity.PostureCheckCreated + var affectedPeerIDs []string err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { if err = validatePostureChecks(ctx, transaction, accountID, postureChecks); err != nil { @@ -50,12 +50,10 @@ func (am *DefaultAccountManager) SavePostureChecks(ctx context.Context, accountI } if isUpdate { - updateAccountPeers, err = arePostureCheckChangesAffectPeers(ctx, transaction, accountID, postureChecks.ID) - if err != nil { - return err - } - action = activity.PostureCheckUpdated + + groupIDs, directPeerIDs := collectPostureCheckAffectedGroupsAndPeers(ctx, transaction, accountID, postureChecks.ID) + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) } postureChecks.AccountID = accountID @@ -75,8 +73,8 @@ func (am *DefaultAccountManager) SavePostureChecks(ctx context.Context, accountI am.StoreEvent(ctx, userID, postureChecks.ID, accountID, action, postureChecks.EventMeta()) - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return postureChecks, nil @@ -132,27 +130,23 @@ func (am *DefaultAccountManager) ListPostureChecks(ctx context.Context, accountI return am.Store.GetAccountPostureChecks(ctx, store.LockingStrengthNone, accountID) } -// arePostureCheckChangesAffectPeers checks if the changes in posture checks are affecting peers. -func arePostureCheckChangesAffectPeers(ctx context.Context, transaction store.Store, accountID, postureCheckID string) (bool, error) { +// collectPostureCheckAffectedGroupsAndPeers finds all policies referencing the given posture check +// and collects their affected group IDs and direct peer IDs. +func collectPostureCheckAffectedGroupsAndPeers(ctx context.Context, transaction store.Store, accountID, postureCheckID string) (groupIDs []string, directPeerIDs []string) { policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) if err != nil { - return false, err + return nil, nil } for _, policy := range policies { if slices.Contains(policy.SourcePostureChecks, postureCheckID) { - hasPeers, err := anyGroupHasPeersOrResources(ctx, transaction, accountID, policy.RuleGroups()) - if err != nil { - return false, err - } - - if hasPeers { - return true, nil - } + gIDs, pIDs := collectPolicyAffectedGroupsAndPeers(policy) + groupIDs = append(groupIDs, gIDs...) + directPeerIDs = append(directPeerIDs, pIDs...) } } - return false, nil + return groupIDs, directPeerIDs } // validatePostureChecks validates the posture checks. diff --git a/management/server/route.go b/management/server/route.go index 2b4f11d05..30297f851 100644 --- a/management/server/route.go +++ b/management/server/route.go @@ -147,7 +147,7 @@ func (am *DefaultAccountManager) CreateRoute(ctx context.Context, accountID stri } var newRoute *route.Route - var updateAccountPeers bool + var affectedPeerIDs []string err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { newRoute = &route.Route{ @@ -173,15 +173,13 @@ func (am *DefaultAccountManager) CreateRoute(ctx context.Context, accountID stri return err } - updateAccountPeers, err = areRouteChangesAffectPeers(ctx, transaction, newRoute) - if err != nil { - return err - } - if err = transaction.SaveRoute(ctx, newRoute); err != nil { return err } + groupIDs, directPeerIDs := collectRouteAffectedGroupsAndPeers(newRoute) + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) + return transaction.IncrementNetworkSerial(ctx, accountID) }) if err != nil { @@ -190,8 +188,8 @@ func (am *DefaultAccountManager) CreateRoute(ctx context.Context, accountID stri am.StoreEvent(ctx, userID, string(newRoute.ID), accountID, activity.RouteCreated, newRoute.EventMeta()) - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return newRoute, nil @@ -208,8 +206,7 @@ func (am *DefaultAccountManager) SaveRoute(ctx context.Context, accountID, userI } var oldRoute *route.Route - var oldRouteAffectsPeers bool - var newRouteAffectsPeers bool + var affectedPeerIDs []string err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { if err = validateRoute(ctx, transaction, accountID, routeToSave); err != nil { @@ -221,21 +218,15 @@ func (am *DefaultAccountManager) SaveRoute(ctx context.Context, accountID, userI return err } - oldRouteAffectsPeers, err = areRouteChangesAffectPeers(ctx, transaction, oldRoute) - if err != nil { - return err - } - - newRouteAffectsPeers, err = areRouteChangesAffectPeers(ctx, transaction, routeToSave) - if err != nil { - return err - } routeToSave.AccountID = accountID if err = transaction.SaveRoute(ctx, routeToSave); err != nil { return err } + groupIDs, directPeerIDs := collectRouteAffectedGroupsAndPeers(routeToSave, oldRoute) + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) + return transaction.IncrementNetworkSerial(ctx, accountID) }) if err != nil { @@ -244,8 +235,8 @@ func (am *DefaultAccountManager) SaveRoute(ctx context.Context, accountID, userI am.StoreEvent(ctx, userID, string(routeToSave.ID), accountID, activity.RouteUpdated, routeToSave.EventMeta()) - if oldRouteAffectsPeers || newRouteAffectsPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return nil @@ -261,19 +252,17 @@ func (am *DefaultAccountManager) DeleteRoute(ctx context.Context, accountID stri return status.NewPermissionDeniedError() } - var route *route.Route - var updateAccountPeers bool + var rt *route.Route + var affectedPeerIDs []string err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { - route, err = transaction.GetRouteByID(ctx, store.LockingStrengthUpdate, accountID, string(routeID)) + rt, err = transaction.GetRouteByID(ctx, store.LockingStrengthUpdate, accountID, string(routeID)) if err != nil { return err } - updateAccountPeers, err = areRouteChangesAffectPeers(ctx, transaction, route) - if err != nil { - return err - } + groupIDs, directPeerIDs := collectRouteAffectedGroupsAndPeers(rt) + affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) if err = transaction.DeleteRoute(ctx, accountID, string(routeID)); err != nil { return err @@ -285,10 +274,10 @@ func (am *DefaultAccountManager) DeleteRoute(ctx context.Context, accountID stri return fmt.Errorf("failed to delete route %s: %w", routeID, err) } - am.StoreEvent(ctx, userID, string(route.ID), accountID, activity.RouteRemoved, route.EventMeta()) + am.StoreEvent(ctx, userID, string(rt.ID), accountID, activity.RouteRemoved, rt.EventMeta()) - if updateAccountPeers { - am.UpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) > 0 { + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } return nil @@ -377,23 +366,20 @@ func getPlaceholderIP() netip.Prefix { return netip.PrefixFrom(netip.AddrFrom4([4]byte{192, 0, 2, 0}), 32) } -// areRouteChangesAffectPeers checks if a given route affects peers by determining -// if it has a routing peer, distribution, or peer groups that include peers. -func areRouteChangesAffectPeers(ctx context.Context, transaction store.Store, route *route.Route) (bool, error) { - if route.Peer != "" { - return true, nil +// collectRouteAffectedGroupsAndPeers returns group IDs and direct peer IDs from the given routes. +func collectRouteAffectedGroupsAndPeers(routes ...*route.Route) (groupIDs []string, directPeerIDs []string) { + for _, r := range routes { + if r == nil { + continue + } + groupIDs = append(groupIDs, r.Groups...) + groupIDs = append(groupIDs, r.PeerGroups...) + groupIDs = append(groupIDs, r.AccessControlGroups...) + if r.Peer != "" { + directPeerIDs = append(directPeerIDs, r.Peer) + } } - - hasPeers, err := anyGroupHasPeersOrResources(ctx, transaction, route.AccountID, route.Groups) - if err != nil { - return false, err - } - - if hasPeers { - return true, nil - } - - return anyGroupHasPeersOrResources(ctx, transaction, route.AccountID, route.PeerGroups) + return } // GetRoutesByPrefixOrDomains return list of routes by account and route prefix diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go index 0a716d08d..c285b70c7 100644 --- a/management/server/store/sql_store.go +++ b/management/server/store/sql_store.go @@ -4662,6 +4662,23 @@ func (s *SqlStore) GetPeersByGroupIDs(ctx context.Context, accountID string, gro return peers, nil } +func (s *SqlStore) GetPeerIDsByGroups(ctx context.Context, accountID string, groupIDs []string) ([]string, error) { + if len(groupIDs) == 0 { + return nil, nil + } + + var peerIDs []string + result := s.db.Model(&types.GroupPeer{}). + Select("DISTINCT peer_id"). + Where("account_id = ? AND group_id IN ?", accountID, groupIDs). + Pluck("peer_id", &peerIDs) + if result.Error != nil { + return nil, status.Errorf(status.Internal, "failed to get peer IDs by groups: %s", result.Error) + } + + return peerIDs, nil +} + func (s *SqlStore) GetUserIDByPeerKey(ctx context.Context, lockStrength LockingStrength, peerKey string) (string, error) { tx := s.db if lockStrength != LockingStrengthNone { diff --git a/management/server/store/store.go b/management/server/store/store.go index 0d8b0678a..82489615f 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -159,6 +159,7 @@ type Store interface { GetPeerByID(ctx context.Context, lockStrength LockingStrength, accountID string, peerID string) (*nbpeer.Peer, error) GetPeersByIDs(ctx context.Context, lockStrength LockingStrength, accountID string, peerIDs []string) (map[string]*nbpeer.Peer, error) GetPeersByGroupIDs(ctx context.Context, accountID string, groupIDs []string) ([]*nbpeer.Peer, error) + GetPeerIDsByGroups(ctx context.Context, accountID string, groupIDs []string) ([]string, error) GetAccountPeersWithExpiration(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*nbpeer.Peer, error) GetAccountPeersWithInactivity(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*nbpeer.Peer, error) GetAllEphemeralPeers(ctx context.Context, lockStrength LockingStrength) ([]*nbpeer.Peer, error) diff --git a/management/server/store/store_mock.go b/management/server/store/store_mock.go index beee13d96..70366ed44 100644 --- a/management/server/store/store_mock.go +++ b/management/server/store/store_mock.go @@ -178,6 +178,7 @@ func (mr *MockStoreMockRecorder) GetClusterSupportsCrowdSec(ctx, clusterAddr int mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetClusterSupportsCrowdSec", reflect.TypeOf((*MockStore)(nil).GetClusterSupportsCrowdSec), ctx, clusterAddr) } + // Close mocks base method. func (m *MockStore) Close(ctx context.Context) error { m.ctrl.T.Helper() @@ -1852,6 +1853,21 @@ func (mr *MockStoreMockRecorder) GetPeersByGroupIDs(ctx, accountID, groupIDs int return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeersByGroupIDs", reflect.TypeOf((*MockStore)(nil).GetPeersByGroupIDs), ctx, accountID, groupIDs) } +// GetPeerIDsByGroups mocks base method. +func (m *MockStore) GetPeerIDsByGroups(ctx context.Context, accountID string, groupIDs []string) ([]string, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetPeerIDsByGroups", ctx, accountID, groupIDs) + ret0, _ := ret[0].([]string) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetPeerIDsByGroups indicates an expected call of GetPeerIDsByGroups. +func (mr *MockStoreMockRecorder) GetPeerIDsByGroups(ctx, accountID, groupIDs interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeerIDsByGroups", reflect.TypeOf((*MockStore)(nil).GetPeerIDsByGroups), ctx, accountID, groupIDs) +} + // GetPeersByIDs mocks base method. func (m *MockStore) GetPeersByIDs(ctx context.Context, lockStrength LockingStrength, accountID string, peerIDs []string) (map[string]*peer.Peer, error) { m.ctrl.T.Helper() From 5a16c812fd729cccf6c21728d34b23a1ffecbd0c Mon Sep 17 00:00:00 2001 From: pascal Date: Mon, 27 Apr 2026 18:18:29 +0200 Subject: [PATCH 02/28] use buffering affected peers --- .../network_map/controller/controller.go | 100 ++++++++++++++++++ .../controllers/network_map/interface.go | 1 + .../controllers/network_map/interface_mock.go | 14 +++ management/server/account/manager.go | 1 + management/server/account/manager_mock.go | 12 +++ management/server/mock_server/account_mock.go | 7 ++ management/server/peer.go | 6 ++ 7 files changed, 141 insertions(+) diff --git a/management/internals/controllers/network_map/controller/controller.go b/management/internals/controllers/network_map/controller/controller.go index 5cbf6fceb..f13eafbcf 100644 --- a/management/internals/controllers/network_map/controller/controller.go +++ b/management/internals/controllers/network_map/controller/controller.go @@ -45,6 +45,7 @@ type Controller struct { accountUpdateLocks sync.Map sendAccountUpdateLocks sync.Map + affectedPeerUpdateLocks sync.Map updateAccountPeersBufferInterval atomic.Int64 // dnsDomain is used for peer resolution. This is appended to the peer's name dnsDomain string @@ -63,6 +64,13 @@ type bufferUpdate struct { update atomic.Bool } +type bufferAffectedUpdate struct { + sendMu sync.Mutex + dataMu sync.Mutex + next *time.Timer + peerIDs map[string]struct{} +} + var _ network_map.Controller = (*Controller)(nil) func NewController(ctx context.Context, store store.Store, metrics telemetry.AppMetrics, peersUpdateManager network_map.PeersUpdateManager, requestBuffer account.RequestBuffer, integratedPeerValidator integrated_validator.IntegratedValidator, settingsManager settings.Manager, dnsDomain string, proxyController port_forwarding.Controller, ephemeralPeersManager ephemeral.Manager, config *config.Config) *Controller { @@ -496,6 +504,98 @@ func (c *Controller) BufferUpdateAccountPeers(ctx context.Context, accountID str return nil } +// BufferUpdateAffectedPeers accumulates peer IDs across rapid successive calls +// and flushes them in a single sendUpdateForAffectedPeers call after the buffer interval. +func (c *Controller) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error { + if len(peerIDs) == 0 { + return nil + } + + log.WithContext(ctx).Tracef("buffer updating %d affected peers for account %s from %s", len(peerIDs), accountID, util.GetCallerName()) + + bufUpd, _ := c.affectedPeerUpdateLocks.LoadOrStore(accountID, &bufferAffectedUpdate{ + peerIDs: make(map[string]struct{}), + }) + b := bufUpd.(*bufferAffectedUpdate) + + // Always accumulate incoming peer IDs (non-blocking). + b.addPeerIDs(peerIDs) + + if !b.sendMu.TryLock() { + // Another goroutine is already sending; it will pick up our IDs on its next drain. + return nil + } + + b.stopTimer() + + collected := b.drainPeerIDs() + go func() { + defer b.sendMu.Unlock() + _ = c.sendUpdateForAffectedPeers(ctx, accountID, collected) + + // Check if more peer IDs accumulated while we were sending. + if !b.hasPending() { + return + } + + // Schedule a debounced flush for the newly accumulated IDs. + b.setTimer(time.Duration(c.updateAccountPeersBufferInterval.Load()), func() { + ids := b.drainPeerIDs() + if len(ids) > 0 { + _ = c.sendUpdateForAffectedPeers(ctx, accountID, ids) + } + }) + }() + + return nil +} + +func (b *bufferAffectedUpdate) addPeerIDs(ids []string) { + b.dataMu.Lock() + for _, id := range ids { + b.peerIDs[id] = struct{}{} + } + b.dataMu.Unlock() +} + +func (b *bufferAffectedUpdate) drainPeerIDs() []string { + b.dataMu.Lock() + defer b.dataMu.Unlock() + if len(b.peerIDs) == 0 { + return nil + } + ids := make([]string, 0, len(b.peerIDs)) + for id := range b.peerIDs { + ids = append(ids, id) + } + b.peerIDs = make(map[string]struct{}) + return ids +} + +func (b *bufferAffectedUpdate) hasPending() bool { + b.dataMu.Lock() + defer b.dataMu.Unlock() + return len(b.peerIDs) > 0 +} + +func (b *bufferAffectedUpdate) stopTimer() { + b.dataMu.Lock() + defer b.dataMu.Unlock() + if b.next != nil { + b.next.Stop() + } +} + +func (b *bufferAffectedUpdate) setTimer(d time.Duration, f func()) { + b.dataMu.Lock() + defer b.dataMu.Unlock() + if b.next == nil { + b.next = time.AfterFunc(d, f) + return + } + b.next.Reset(d) +} + func (c *Controller) GetValidatedPeerWithMap(ctx context.Context, isRequiresApproval bool, accountID string, peer *nbpeer.Peer) (*nbpeer.Peer, *types.NetworkMap, []*posture.Checks, int64, error) { if isRequiresApproval { network, err := c.repo.GetAccountNetwork(ctx, accountID) diff --git a/management/internals/controllers/network_map/interface.go b/management/internals/controllers/network_map/interface.go index 8d81556f9..4b9fdee12 100644 --- a/management/internals/controllers/network_map/interface.go +++ b/management/internals/controllers/network_map/interface.go @@ -20,6 +20,7 @@ const ( type Controller interface { UpdateAccountPeers(ctx context.Context, accountID string) error UpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error + BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error UpdateAccountPeer(ctx context.Context, accountId string, peerId string) error BufferUpdateAccountPeers(ctx context.Context, accountID string) error GetValidatedPeerWithMap(ctx context.Context, isRequiresApproval bool, accountID string, p *nbpeer.Peer) (*nbpeer.Peer, *types.NetworkMap, []*posture.Checks, int64, error) diff --git a/management/internals/controllers/network_map/interface_mock.go b/management/internals/controllers/network_map/interface_mock.go index b2ef0b861..15b6bdc56 100644 --- a/management/internals/controllers/network_map/interface_mock.go +++ b/management/internals/controllers/network_map/interface_mock.go @@ -57,6 +57,20 @@ func (mr *MockControllerMockRecorder) BufferUpdateAccountPeers(ctx, accountID an return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "BufferUpdateAccountPeers", reflect.TypeOf((*MockController)(nil).BufferUpdateAccountPeers), ctx, accountID) } +// BufferUpdateAffectedPeers mocks base method. +func (m *MockController) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "BufferUpdateAffectedPeers", ctx, accountID, peerIDs) + ret0, _ := ret[0].(error) + return ret0 +} + +// BufferUpdateAffectedPeers indicates an expected call of BufferUpdateAffectedPeers. +func (mr *MockControllerMockRecorder) BufferUpdateAffectedPeers(ctx, accountID, peerIDs any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "BufferUpdateAffectedPeers", reflect.TypeOf((*MockController)(nil).BufferUpdateAffectedPeers), ctx, accountID, peerIDs) +} + // CountStreams mocks base method. func (m *MockController) CountStreams() int { m.ctrl.T.Helper() diff --git a/management/server/account/manager.go b/management/server/account/manager.go index 576054d1e..9b4a8152c 100644 --- a/management/server/account/manager.go +++ b/management/server/account/manager.go @@ -126,6 +126,7 @@ type Manager interface { DeleteSetupKey(ctx context.Context, accountID, userID, keyID string) error UpdateAccountPeers(ctx context.Context, accountID string) UpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) + BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) BufferUpdateAccountPeers(ctx context.Context, accountID string) BuildUserInfosForAccount(ctx context.Context, accountID, initiatorUserID string, accountUsers []*types.User) (map[string]*types.UserInfo, error) SyncUserJWTGroups(ctx context.Context, userAuth auth.UserAuth) error diff --git a/management/server/account/manager_mock.go b/management/server/account/manager_mock.go index c595346f8..0896255ac 100644 --- a/management/server/account/manager_mock.go +++ b/management/server/account/manager_mock.go @@ -122,6 +122,18 @@ func (mr *MockManagerMockRecorder) BufferUpdateAccountPeers(ctx, accountID inter return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "BufferUpdateAccountPeers", reflect.TypeOf((*MockManager)(nil).BufferUpdateAccountPeers), ctx, accountID) } +// BufferUpdateAffectedPeers mocks base method. +func (m *MockManager) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "BufferUpdateAffectedPeers", ctx, accountID, peerIDs) +} + +// BufferUpdateAffectedPeers indicates an expected call of BufferUpdateAffectedPeers. +func (mr *MockManagerMockRecorder) BufferUpdateAffectedPeers(ctx, accountID, peerIDs interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "BufferUpdateAffectedPeers", reflect.TypeOf((*MockManager)(nil).BufferUpdateAffectedPeers), ctx, accountID, peerIDs) +} + // BuildUserInfosForAccount mocks base method. func (m *MockManager) BuildUserInfosForAccount(ctx context.Context, accountID, initiatorUserID string, accountUsers []*types.User) (map[string]*types.UserInfo, error) { m.ctrl.T.Helper() diff --git a/management/server/mock_server/account_mock.go b/management/server/mock_server/account_mock.go index 5a9009da7..d4580a4c6 100644 --- a/management/server/mock_server/account_mock.go +++ b/management/server/mock_server/account_mock.go @@ -130,6 +130,7 @@ type MockAccountManager struct { AllowSyncFunc func(string, uint64) bool UpdateAccountPeersFunc func(ctx context.Context, accountID string) UpdateAffectedPeersFunc func(ctx context.Context, accountID string, peerIDs []string) + BufferUpdateAffectedPeersFunc func(ctx context.Context, accountID string, peerIDs []string) BufferUpdateAccountPeersFunc func(ctx context.Context, accountID string) RecalculateNetworkMapCacheFunc func(ctx context.Context, accountId string) error @@ -213,6 +214,12 @@ func (am *MockAccountManager) UpdateAffectedPeers(ctx context.Context, accountID } } +func (am *MockAccountManager) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) { + if am.BufferUpdateAffectedPeersFunc != nil { + am.BufferUpdateAffectedPeersFunc(ctx, accountID, peerIDs) + } +} + func (am *MockAccountManager) BufferUpdateAccountPeers(ctx context.Context, accountID string) { if am.BufferUpdateAccountPeersFunc != nil { am.BufferUpdateAccountPeersFunc(ctx, accountID) diff --git a/management/server/peer.go b/management/server/peer.go index 39368c840..509d53ebb 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -1326,6 +1326,12 @@ func (am *DefaultAccountManager) resolvePeerIDs(ctx context.Context, s store.Sto return peerIDs } +// BufferUpdateAffectedPeers accumulates peer IDs across rapid successive calls +// and flushes them in a single update after the buffer interval. +func (am *DefaultAccountManager) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) { + _ = am.networkMapController.BufferUpdateAffectedPeers(ctx, accountID, peerIDs) +} + func (am *DefaultAccountManager) BufferUpdateAccountPeers(ctx context.Context, accountID string) { _ = am.networkMapController.BufferUpdateAccountPeers(ctx, accountID) } From cefb37e920b9153837b46c13c8c10e3370f98544 Mon Sep 17 00:00:00 2001 From: pascal Date: Tue, 28 Apr 2026 13:48:01 +0200 Subject: [PATCH 03/28] affected filtering on peers update --- .../network_map/controller/controller.go | 27 +++++++---- .../controllers/network_map/interface.go | 6 +-- .../controllers/network_map/interface_mock.go | 24 +++++----- management/server/account.go | 4 +- management/server/group.go | 17 +++++-- management/server/peer.go | 48 ++++++++++++++++--- management/server/store/sql_store.go | 17 +++++++ management/server/store/store.go | 1 + management/server/store/store_mock.go | 15 ++++++ management/server/user.go | 14 +++++- management/server/user_test.go | 6 +-- 11 files changed, 138 insertions(+), 41 deletions(-) diff --git a/management/internals/controllers/network_map/controller/controller.go b/management/internals/controllers/network_map/controller/controller.go index f13eafbcf..b14d0d81a 100644 --- a/management/internals/controllers/network_map/controller/controller.go +++ b/management/internals/controllers/network_map/controller/controller.go @@ -793,21 +793,24 @@ func isPeerInPolicySourceGroups(account *types.Account, peerID string, policy *t return false, nil } -func (c *Controller) OnPeersUpdated(ctx context.Context, accountID string, peerIDs []string) error { - err := c.bufferSendUpdateAccountPeers(ctx, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to buffer update account peers for peer update in account %s: %v", accountID, err) +func (c *Controller) OnPeersUpdated(ctx context.Context, accountID string, peerIDs []string, affectedPeerIDs []string) error { + if len(affectedPeerIDs) == 0 { + log.WithContext(ctx).Tracef("no affected peers for peer update in account %s, skipping", accountID) + return nil } - - return nil + return c.BufferUpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } -func (c *Controller) OnPeersAdded(ctx context.Context, accountID string, peerIDs []string) error { +func (c *Controller) OnPeersAdded(ctx context.Context, accountID string, peerIDs []string, affectedPeerIDs []string) error { log.WithContext(ctx).Debugf("OnPeersAdded call to add peers: %v", peerIDs) - return c.bufferSendUpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) == 0 { + log.WithContext(ctx).Tracef("no affected peers for peer add in account %s, skipping", accountID) + return nil + } + return c.BufferUpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } -func (c *Controller) OnPeersDeleted(ctx context.Context, accountID string, peerIDs []string) error { +func (c *Controller) OnPeersDeleted(ctx context.Context, accountID string, peerIDs []string, affectedPeerIDs []string) error { network, err := c.repo.GetAccountNetwork(ctx, accountID) if err != nil { return err @@ -840,7 +843,11 @@ func (c *Controller) OnPeersDeleted(ctx context.Context, accountID string, peerI c.peersUpdateManager.CloseChannel(ctx, peerID) } - return c.bufferSendUpdateAccountPeers(ctx, accountID) + if len(affectedPeerIDs) == 0 { + log.WithContext(ctx).Tracef("no affected peers for peer delete in account %s, skipping network map update", accountID) + return nil + } + return c.BufferUpdateAffectedPeers(ctx, accountID, affectedPeerIDs) } // GetNetworkMap returns Network map for a given peer (omits original peer from the Peers result) diff --git a/management/internals/controllers/network_map/interface.go b/management/internals/controllers/network_map/interface.go index 4b9fdee12..8cd84b605 100644 --- a/management/internals/controllers/network_map/interface.go +++ b/management/internals/controllers/network_map/interface.go @@ -29,9 +29,9 @@ type Controller interface { GetNetworkMap(ctx context.Context, peerID string) (*types.NetworkMap, error) CountStreams() int - OnPeersUpdated(ctx context.Context, accountId string, peerIDs []string) error - OnPeersAdded(ctx context.Context, accountID string, peerIDs []string) error - OnPeersDeleted(ctx context.Context, accountID string, peerIDs []string) error + OnPeersUpdated(ctx context.Context, accountId string, peerIDs []string, affectedPeerIDs []string) error + OnPeersAdded(ctx context.Context, accountID string, peerIDs []string, affectedPeerIDs []string) error + 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) diff --git a/management/internals/controllers/network_map/interface_mock.go b/management/internals/controllers/network_map/interface_mock.go index 15b6bdc56..1178969ff 100644 --- a/management/internals/controllers/network_map/interface_mock.go +++ b/management/internals/controllers/network_map/interface_mock.go @@ -172,45 +172,45 @@ func (mr *MockControllerMockRecorder) OnPeerDisconnected(ctx, accountID, peerID } // OnPeersAdded mocks base method. -func (m *MockController) OnPeersAdded(ctx context.Context, accountID string, peerIDs []string) error { +func (m *MockController) OnPeersAdded(ctx context.Context, accountID string, peerIDs []string, affectedPeerIDs []string) error { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "OnPeersAdded", ctx, accountID, peerIDs) + ret := m.ctrl.Call(m, "OnPeersAdded", ctx, accountID, peerIDs, affectedPeerIDs) ret0, _ := ret[0].(error) return ret0 } // OnPeersAdded indicates an expected call of OnPeersAdded. -func (mr *MockControllerMockRecorder) OnPeersAdded(ctx, accountID, peerIDs any) *gomock.Call { +func (mr *MockControllerMockRecorder) OnPeersAdded(ctx, accountID, peerIDs, affectedPeerIDs any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "OnPeersAdded", reflect.TypeOf((*MockController)(nil).OnPeersAdded), ctx, accountID, peerIDs) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "OnPeersAdded", reflect.TypeOf((*MockController)(nil).OnPeersAdded), ctx, accountID, peerIDs, affectedPeerIDs) } // OnPeersDeleted mocks base method. -func (m *MockController) OnPeersDeleted(ctx context.Context, accountID string, peerIDs []string) error { +func (m *MockController) OnPeersDeleted(ctx context.Context, accountID string, peerIDs []string, affectedPeerIDs []string) error { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "OnPeersDeleted", ctx, accountID, peerIDs) + ret := m.ctrl.Call(m, "OnPeersDeleted", ctx, accountID, peerIDs, affectedPeerIDs) ret0, _ := ret[0].(error) return ret0 } // OnPeersDeleted indicates an expected call of OnPeersDeleted. -func (mr *MockControllerMockRecorder) OnPeersDeleted(ctx, accountID, peerIDs any) *gomock.Call { +func (mr *MockControllerMockRecorder) OnPeersDeleted(ctx, accountID, peerIDs, affectedPeerIDs any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "OnPeersDeleted", reflect.TypeOf((*MockController)(nil).OnPeersDeleted), ctx, accountID, peerIDs) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "OnPeersDeleted", reflect.TypeOf((*MockController)(nil).OnPeersDeleted), ctx, accountID, peerIDs, affectedPeerIDs) } // OnPeersUpdated mocks base method. -func (m *MockController) OnPeersUpdated(ctx context.Context, accountId string, peerIDs []string) error { +func (m *MockController) OnPeersUpdated(ctx context.Context, accountId string, peerIDs []string, affectedPeerIDs []string) error { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "OnPeersUpdated", ctx, accountId, peerIDs) + ret := m.ctrl.Call(m, "OnPeersUpdated", ctx, accountId, peerIDs, affectedPeerIDs) ret0, _ := ret[0].(error) return ret0 } // OnPeersUpdated indicates an expected call of OnPeersUpdated. -func (mr *MockControllerMockRecorder) OnPeersUpdated(ctx, accountId, peerIDs any) *gomock.Call { +func (mr *MockControllerMockRecorder) OnPeersUpdated(ctx, accountId, peerIDs, affectedPeerIDs any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "OnPeersUpdated", reflect.TypeOf((*MockController)(nil).OnPeersUpdated), ctx, accountId, peerIDs) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "OnPeersUpdated", reflect.TypeOf((*MockController)(nil).OnPeersUpdated), ctx, accountId, peerIDs, affectedPeerIDs) } // StartWarmup mocks base method. diff --git a/management/server/account.go b/management/server/account.go index 7d53cef03..cbf67904c 100644 --- a/management/server/account.go +++ b/management/server/account.go @@ -2215,7 +2215,9 @@ func (am *DefaultAccountManager) UpdatePeerIP(ctx context.Context, accountID, us if err != nil { return err } - err = am.networkMapController.OnPeersUpdated(ctx, peer.AccountID, []string{peerID}) + changedPeerIDs := []string{peerID} + affectedPeerIDs := am.resolveAffectedPeersForPeerChanges(ctx, am.Store, accountID, changedPeerIDs) + err = am.networkMapController.OnPeersUpdated(ctx, peer.AccountID, changedPeerIDs, affectedPeerIDs) if err != nil { return fmt.Errorf("notify network map controller of peer update: %w", err) } diff --git a/management/server/group.go b/management/server/group.go index 4bd249398..cc19cf1a4 100644 --- a/management/server/group.go +++ b/management/server/group.go @@ -835,11 +835,9 @@ func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Sto changedSet[id] = struct{}{} } + log.WithContext(ctx).Tracef("collecting affected groups for changed groups %v", changedGroupIDs) + groupSet := make(map[string]struct{}) - // Always include the changed groups themselves - for _, id := range changedGroupIDs { - groupSet[id] = struct{}{} - } peerSet := make(map[string]struct{}) @@ -852,14 +850,17 @@ func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Sto if !policyReferencesGroups(policy, changedSet) { continue } + log.WithContext(ctx).Tracef("policy %s (%s) references changed groups, adding rule groups", policy.ID, policy.Name) for _, gID := range policy.RuleGroups() { groupSet[gID] = struct{}{} } for _, rule := range policy.Rules { if rule.SourceResource.Type == types.ResourceTypePeer && rule.SourceResource.ID != "" { + log.WithContext(ctx).Tracef("policy %s rule %s has direct source peer %s", policy.ID, rule.ID, rule.SourceResource.ID) peerSet[rule.SourceResource.ID] = struct{}{} } if rule.DestinationResource.Type == types.ResourceTypePeer && rule.DestinationResource.ID != "" { + log.WithContext(ctx).Tracef("policy %s rule %s has direct destination peer %s", policy.ID, rule.ID, rule.DestinationResource.ID) peerSet[rule.DestinationResource.ID] = struct{}{} } } @@ -875,6 +876,7 @@ func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Sto if !routeReferencesGroups(r, changedSet) { continue } + log.WithContext(ctx).Tracef("route %s (%s) references changed groups", r.ID, r.Description) for _, gID := range r.Groups { groupSet[gID] = struct{}{} } @@ -885,6 +887,7 @@ func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Sto groupSet[gID] = struct{}{} } if r.Peer != "" { + log.WithContext(ctx).Tracef("route %s has direct peer %s", r.ID, r.Peer) peerSet[r.Peer] = struct{}{} } } @@ -898,6 +901,7 @@ func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Sto for _, ns := range nsGroups { for _, gID := range ns.Groups { if _, ok := changedSet[gID]; ok { + log.WithContext(ctx).Tracef("nameserver group %s (%s) references changed group %s", ns.ID, ns.Name, gID) for _, g := range ns.Groups { groupSet[g] = struct{}{} } @@ -914,6 +918,7 @@ func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Sto } else { for _, gID := range dnsSettings.DisabledManagementGroups { if _, ok := changedSet[gID]; ok { + log.WithContext(ctx).Tracef("DNS disabled management group %s matches changed group", gID) groupSet[gID] = struct{}{} } } @@ -928,10 +933,12 @@ func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Sto if !routerReferencesGroups(router, changedSet) { continue } + log.WithContext(ctx).Tracef("network router %s references changed groups", router.ID) for _, gID := range router.PeerGroups { groupSet[gID] = struct{}{} } if router.Peer != "" { + log.WithContext(ctx).Tracef("network router %s has direct peer %s", router.ID, router.Peer) peerSet[router.Peer] = struct{}{} } } @@ -947,6 +954,8 @@ func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Sto directPeerIDs = append(directPeerIDs, pID) } + log.WithContext(ctx).Tracef("affected groups resolution: changed=%v -> affectedGroups=%v, directPeers=%v", changedGroupIDs, allGroupIDs, directPeerIDs) + return allGroupIDs, directPeerIDs } diff --git a/management/server/peer.go b/management/server/peer.go index 509d53ebb..ee3a6369f 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -150,7 +150,9 @@ func (am *DefaultAccountManager) MarkPeerConnected(ctx context.Context, peerPubK } if expired { - err = am.networkMapController.OnPeersUpdated(ctx, accountID, []string{peer.ID}) + changedPeerIDs := []string{peer.ID} + affectedPeerIDs := am.resolveAffectedPeersForPeerChanges(ctx, am.Store, accountID, changedPeerIDs) + err = am.networkMapController.OnPeersUpdated(ctx, accountID, changedPeerIDs, affectedPeerIDs) if err != nil { return fmt.Errorf("notify network map controller of peer update: %w", err) } @@ -334,7 +336,9 @@ func (am *DefaultAccountManager) UpdatePeer(ctx context.Context, accountID, user } } - err = am.networkMapController.OnPeersUpdated(ctx, accountID, []string{peer.ID}) + changedPeerIDs := []string{peer.ID} + affectedPeerIDs := am.resolveAffectedPeersForPeerChanges(ctx, am.Store, accountID, changedPeerIDs) + err = am.networkMapController.OnPeersUpdated(ctx, accountID, changedPeerIDs, affectedPeerIDs) if err != nil { return nil, fmt.Errorf("notify network map controller of peer update: %w", err) } @@ -492,6 +496,7 @@ func (am *DefaultAccountManager) DeletePeer(ctx context.Context, accountID, peer var peer *nbpeer.Peer var settings *types.Settings var eventsToStore []func() + var affectedPeerIDs []string serviceID, err := am.serviceManager.GetServiceIDByTargetID(ctx, accountID, peerID) if err != nil { @@ -516,6 +521,8 @@ func (am *DefaultAccountManager) DeletePeer(ctx context.Context, accountID, peer return err } + affectedPeerIDs = am.resolveAffectedPeersForPeerChanges(ctx, transaction, accountID, []string{peerID}) + eventsToStore, err = deletePeers(ctx, am, transaction, accountID, userID, []*nbpeer.Peer{peer}, settings) if err != nil { return fmt.Errorf("failed to delete peer: %w", err) @@ -539,7 +546,7 @@ func (am *DefaultAccountManager) DeletePeer(ctx context.Context, accountID, peer log.WithContext(ctx).Errorf("failed to delete peer %s from integrated validator: %v", peerID, err) } - if err = am.networkMapController.OnPeersDeleted(ctx, accountID, []string{peerID}); err != nil { + if err = am.networkMapController.OnPeersDeleted(ctx, accountID, []string{peerID}, affectedPeerIDs); err != nil { log.WithContext(ctx).Errorf("failed to delete peer %s from network map: %v", peerID, err) } @@ -863,7 +870,9 @@ func (am *DefaultAccountManager) AddPeer(ctx context.Context, accountID, setupKe am.StoreEvent(ctx, opEvent.InitiatorID, opEvent.TargetID, opEvent.AccountID, opEvent.Activity, opEvent.Meta) } - if err := am.networkMapController.OnPeersAdded(ctx, accountID, []string{newPeer.ID}); err != nil { + changedPeerIDs := []string{newPeer.ID} + affectedPeerIDs := am.resolveAffectedPeersForPeerChanges(ctx, am.Store, accountID, changedPeerIDs) + if err := am.networkMapController.OnPeersAdded(ctx, accountID, changedPeerIDs, affectedPeerIDs); err != nil { log.WithContext(ctx).Errorf("failed to update network map cache for peer %s: %v", newPeer.ID, err) } @@ -946,7 +955,9 @@ func (am *DefaultAccountManager) SyncPeer(ctx context.Context, sync types.PeerSy } if isStatusChanged || sync.UpdateAccountPeers || (updated && (len(postureChecks) > 0 || versionChanged)) { - err = am.networkMapController.OnPeersUpdated(ctx, accountID, []string{peer.ID}) + changedPeerIDs := []string{peer.ID} + affectedPeerIDs := am.resolveAffectedPeersForPeerChanges(ctx, am.Store, accountID, changedPeerIDs) + err = am.networkMapController.OnPeersUpdated(ctx, accountID, changedPeerIDs, affectedPeerIDs) if err != nil { return nil, nil, nil, 0, fmt.Errorf("notify network map controller of peer update: %w", err) } @@ -1073,7 +1084,9 @@ func (am *DefaultAccountManager) LoginPeer(ctx context.Context, login types.Peer } if updateRemotePeers || isStatusChanged || (isPeerUpdated && len(postureChecks) > 0) { - err = am.networkMapController.OnPeersUpdated(ctx, accountID, []string{peer.ID}) + changedPeerIDs := []string{peer.ID} + affectedPeerIDs := am.resolveAffectedPeersForPeerChanges(ctx, am.Store, accountID, changedPeerIDs) + err = am.networkMapController.OnPeersUpdated(ctx, accountID, changedPeerIDs, affectedPeerIDs) if err != nil { return nil, nil, nil, fmt.Errorf("notify network map controller of peer update: %w", err) } @@ -1297,6 +1310,7 @@ func (am *DefaultAccountManager) UpdateAccountPeers(ctx context.Context, account // UpdateAffectedPeers updates only the specified peers that belong to an account. // Should be called when a change is known to affect only a subset of peers. func (am *DefaultAccountManager) UpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) { + log.WithContext(ctx).Tracef("UpdateAffectedPeers: %d peers for account %s", len(peerIDs), accountID) _ = am.networkMapController.UpdateAffectedPeers(ctx, accountID, peerIDs) } @@ -1310,6 +1324,7 @@ func (am *DefaultAccountManager) resolvePeerIDs(ctx context.Context, s store.Sto } if len(directPeerIDs) == 0 { + log.WithContext(ctx).Tracef("resolvePeerIDs: groups=%v -> %d peers", groupIDs, len(peerIDs)) return peerIDs } @@ -1323,6 +1338,8 @@ func (am *DefaultAccountManager) resolvePeerIDs(ctx context.Context, s store.Sto seen[id] = struct{}{} } } + + log.WithContext(ctx).Tracef("resolvePeerIDs: groups=%v + directPeers=%v -> %d peers", groupIDs, directPeerIDs, len(peerIDs)) return peerIDs } @@ -1332,6 +1349,25 @@ func (am *DefaultAccountManager) BufferUpdateAffectedPeers(ctx context.Context, _ = am.networkMapController.BufferUpdateAffectedPeers(ctx, accountID, peerIDs) } +// resolveAffectedPeersForPeerChanges resolves changed peer IDs into the full set of +// affected peers: finds groups containing the changed peers, walks all entity linkages, +// and resolves back to peer IDs. +func (am *DefaultAccountManager) resolveAffectedPeersForPeerChanges(ctx context.Context, s store.Store, accountID string, changedPeerIDs []string) []string { + groupIDs, err := s.GetGroupIDsByPeerIDs(ctx, accountID, changedPeerIDs) + if err != nil { + log.WithContext(ctx).Errorf("failed to get group IDs for changed peers: %v", err) + return nil + } + + log.WithContext(ctx).Tracef("resolveAffectedPeersForPeerChanges: changedPeers=%v -> groups=%v", changedPeerIDs, groupIDs) + + allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, s, accountID, groupIDs) + result := am.resolvePeerIDs(ctx, s, accountID, allGroupIDs, directPeerIDs) + + log.WithContext(ctx).Tracef("resolveAffectedPeersForPeerChanges: changedPeers=%v -> %d affected peers", changedPeerIDs, len(result)) + return result +} + func (am *DefaultAccountManager) BufferUpdateAccountPeers(ctx context.Context, accountID string) { _ = am.networkMapController.BufferUpdateAccountPeers(ctx, accountID) } diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go index c285b70c7..0a2e31538 100644 --- a/management/server/store/sql_store.go +++ b/management/server/store/sql_store.go @@ -4679,6 +4679,23 @@ func (s *SqlStore) GetPeerIDsByGroups(ctx context.Context, accountID string, gro return peerIDs, nil } +func (s *SqlStore) GetGroupIDsByPeerIDs(ctx context.Context, accountID string, peerIDs []string) ([]string, error) { + if len(peerIDs) == 0 { + return nil, nil + } + + var groupIDs []string + result := s.db.Model(&types.GroupPeer{}). + Select("DISTINCT group_id"). + Where("account_id = ? AND peer_id IN ?", accountID, peerIDs). + Pluck("group_id", &groupIDs) + if result.Error != nil { + return nil, status.Errorf(status.Internal, "failed to get group IDs by peers: %s", result.Error) + } + + return groupIDs, nil +} + func (s *SqlStore) GetUserIDByPeerKey(ctx context.Context, lockStrength LockingStrength, peerKey string) (string, error) { tx := s.db if lockStrength != LockingStrengthNone { diff --git a/management/server/store/store.go b/management/server/store/store.go index 82489615f..fddc17f1c 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -160,6 +160,7 @@ type Store interface { GetPeersByIDs(ctx context.Context, lockStrength LockingStrength, accountID string, peerIDs []string) (map[string]*nbpeer.Peer, error) GetPeersByGroupIDs(ctx context.Context, accountID string, groupIDs []string) ([]*nbpeer.Peer, error) GetPeerIDsByGroups(ctx context.Context, accountID string, groupIDs []string) ([]string, error) + GetGroupIDsByPeerIDs(ctx context.Context, accountID string, peerIDs []string) ([]string, error) GetAccountPeersWithExpiration(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*nbpeer.Peer, error) GetAccountPeersWithInactivity(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*nbpeer.Peer, error) GetAllEphemeralPeers(ctx context.Context, lockStrength LockingStrength) ([]*nbpeer.Peer, error) diff --git a/management/server/store/store_mock.go b/management/server/store/store_mock.go index 70366ed44..dac042a18 100644 --- a/management/server/store/store_mock.go +++ b/management/server/store/store_mock.go @@ -1868,6 +1868,21 @@ func (mr *MockStoreMockRecorder) GetPeerIDsByGroups(ctx, accountID, groupIDs int return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeerIDsByGroups", reflect.TypeOf((*MockStore)(nil).GetPeerIDsByGroups), ctx, accountID, groupIDs) } +// GetGroupIDsByPeerIDs mocks base method. +func (m *MockStore) GetGroupIDsByPeerIDs(ctx context.Context, accountID string, peerIDs []string) ([]string, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetGroupIDsByPeerIDs", ctx, accountID, peerIDs) + ret0, _ := ret[0].([]string) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetGroupIDsByPeerIDs indicates an expected call of GetGroupIDsByPeerIDs. +func (mr *MockStoreMockRecorder) GetGroupIDsByPeerIDs(ctx, accountID, peerIDs interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetGroupIDsByPeerIDs", reflect.TypeOf((*MockStore)(nil).GetGroupIDsByPeerIDs), ctx, accountID, peerIDs) +} + // GetPeersByIDs mocks base method. func (m *MockStore) GetPeersByIDs(ctx context.Context, lockStrength LockingStrength, accountID string, peerIDs []string) (map[string]*peer.Peer, error) { m.ctrl.T.Helper() diff --git a/management/server/user.go b/management/server/user.go index c1f984f2f..52a4eadc8 100644 --- a/management/server/user.go +++ b/management/server/user.go @@ -1154,7 +1154,8 @@ func (am *DefaultAccountManager) expireAndUpdatePeers(ctx context.Context, accou } } - err = am.networkMapController.OnPeersUpdated(ctx, accountID, peerIDs) + affectedPeerIDs := am.resolveAffectedPeersForPeerChanges(ctx, am.Store, accountID, peerIDs) + err = am.networkMapController.OnPeersUpdated(ctx, accountID, peerIDs, affectedPeerIDs) if err != nil { return fmt.Errorf("notify network map controller of peer update: %w", err) } @@ -1270,6 +1271,7 @@ func (am *DefaultAccountManager) deleteRegularUser(ctx context.Context, accountI var userPeers []*nbpeer.Peer var targetUser *types.User var settings *types.Settings + var affectedPeerIDs []string var err error err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { @@ -1290,6 +1292,14 @@ func (am *DefaultAccountManager) deleteRegularUser(ctx context.Context, accountI if len(userPeers) > 0 { updateAccountPeers = true + + var peerIDs []string + for _, peer := range userPeers { + peerIDs = append(peerIDs, peer.ID) + } + // Resolve before delete so group memberships are still present. + affectedPeerIDs = am.resolveAffectedPeersForPeerChanges(ctx, transaction, accountID, peerIDs) + addPeerRemovedEvents, err = deletePeers(ctx, am, transaction, accountID, targetUserInfo.ID, userPeers, settings) if err != nil { return fmt.Errorf("failed to delete user peers: %w", err) @@ -1313,7 +1323,7 @@ func (am *DefaultAccountManager) deleteRegularUser(ctx context.Context, accountI log.WithContext(ctx).Errorf("failed to delete peer %s from integrated validator: %v", peer.ID, err) } } - if err := am.networkMapController.OnPeersDeleted(ctx, accountID, peerIDs); err != nil { + if err := am.networkMapController.OnPeersDeleted(ctx, accountID, peerIDs, affectedPeerIDs); err != nil { log.WithContext(ctx).Errorf("failed to delete peers %s from network map: %v", peerIDs, err) } diff --git a/management/server/user_test.go b/management/server/user_test.go index c77ea53d1..68fc58eef 100644 --- a/management/server/user_test.go +++ b/management/server/user_test.go @@ -846,7 +846,7 @@ func TestUser_DeleteUser_regularUser(t *testing.T) { ctrl := gomock.NewController(t) networkMapControllerMock := network_map.NewMockController(ctrl) networkMapControllerMock.EXPECT(). - OnPeersDeleted(gomock.Any(), gomock.Any(), gomock.Any()). + OnPeersDeleted(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()). Return(nil) permissionsManager := permissions.NewManager(store) @@ -962,7 +962,7 @@ func TestUser_DeleteUser_RegularUsers(t *testing.T) { ctrl := gomock.NewController(t) networkMapControllerMock := network_map.NewMockController(ctrl) networkMapControllerMock.EXPECT(). - OnPeersDeleted(gomock.Any(), gomock.Any(), gomock.Any()). + OnPeersDeleted(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()). Return(nil). AnyTimes() @@ -2022,7 +2022,7 @@ func TestUser_Operations_WithEmbeddedIDP(t *testing.T) { ctrl := gomock.NewController(t) networkMapControllerMock := network_map.NewMockController(ctrl) networkMapControllerMock.EXPECT(). - OnPeersDeleted(gomock.Any(), gomock.Any(), gomock.Any()). + OnPeersDeleted(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()). Return(nil). AnyTimes() From 5ec8bebfa5edcea2e545a12c2491857b5c589bbd Mon Sep 17 00:00:00 2001 From: pascal Date: Tue, 28 Apr 2026 16:27:44 +0200 Subject: [PATCH 04/28] add tests --- management/server/affected_peers_test.go | 1159 ++++++++++++++++++++++ management/server/posture_checks_test.go | 41 +- 2 files changed, 1179 insertions(+), 21 deletions(-) create mode 100644 management/server/affected_peers_test.go diff --git a/management/server/affected_peers_test.go b/management/server/affected_peers_test.go new file mode 100644 index 000000000..cf0cb033d --- /dev/null +++ b/management/server/affected_peers_test.go @@ -0,0 +1,1159 @@ +package server + +import ( + "context" + "fmt" + "net/netip" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" + + nbdns "github.com/netbirdio/netbird/dns" + routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" + networkTypes "github.com/netbirdio/netbird/management/server/networks/types" + nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/route" +) + +// setupAffectedPeersTest creates a manager with a clean account and 5 peers, each in its own group. +// Returns the manager, store, account ID, peer IDs, and group IDs. +// Peer layout: +// +// peer0 -> group0 +// peer1 -> group1 +// peer2 -> group2 +// peer3 -> group3 +// peer4 -> group4 +func setupAffectedPeersTest(t *testing.T) (*DefaultAccountManager, store.Store, string, []string, []string) { + t.Helper() + + manager, _, err := createManager(t) + require.NoError(t, err) + + account, err := createAccount(manager, "affected_test", userID, "") + require.NoError(t, err) + + ctx := context.Background() + accountID := account.Id + + // Delete the default "All <-> All" policy so tests start with a clean slate + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + err := manager.Store.DeletePolicy(ctx, accountID, p.ID) + require.NoError(t, err) + } + + setupKey, err := manager.CreateSetupKey(ctx, accountID, "test-key", types.SetupKeyReusable, time.Hour, nil, 999, userID, false, false) + require.NoError(t, err) + + peerIDs := make([]string, 5) + for i := 0; i < 5; i++ { + peer := addPeerToAccount(t, manager, accountID, setupKey.Key) + peerIDs[i] = peer.ID + } + + groupIDs := make([]string, 5) + for i := 0; i < 5; i++ { + g := &types.Group{ + ID: affectedGroupID(i), + Name: affectedGroupName(i), + Peers: []string{peerIDs[i]}, + } + err := manager.CreateGroup(ctx, accountID, userID, g) + require.NoError(t, err) + groupIDs[i] = g.ID + } + + return manager, manager.Store, accountID, peerIDs, groupIDs +} + +func affectedGroupID(i int) string { return fmt.Sprintf("affected-grp-%d", i) } +func affectedGroupName(i int) string { return fmt.Sprintf("AffectedGroup%d", i) } + +// ---------- collectGroupChangeAffectedGroups ---------- + +func TestCollectGroupChange_NoEntities(t *testing.T) { + _, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // group0 is not referenced by any entity + groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) + assert.Empty(t, groups, "no entities reference group0, should return no affected groups") + assert.Empty(t, directPeers, "no entities reference group0, should return no direct peers") +} + +func TestCollectGroupChange_EmptyInput(t *testing.T) { + _, s, accountID, _, _ := setupAffectedPeersTest(t) + ctx := context.Background() + + groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, nil) + assert.Nil(t, groups) + assert.Nil(t, directPeers) + + groups, directPeers = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{}) + assert.Nil(t, groups) + assert.Nil(t, directPeers) +} + +func TestCollectGroupChange_PolicyLinked(t *testing.T) { + manager, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Create policy: group0 (src) <-> group1 (dst) + _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[0]}, + Destinations: []string{groupIDs[1]}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + // Changing group0 should include both group0 and group1 (from the policy rule) + groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) + assert.Contains(t, groups, groupIDs[0]) + assert.Contains(t, groups, groupIDs[1]) + + // Changing group1 should also include both groups + groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) + assert.Contains(t, groups, groupIDs[0]) + assert.Contains(t, groups, groupIDs[1]) + + // Changing group2 (not in policy) should return nothing + groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[2]}) + assert.Empty(t, groups) +} + +func TestCollectGroupChange_PolicyWithDirectPeerResource(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Create policy with direct peer resource + _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[0]}, + SourceResource: types.Resource{ID: peerIDs[3], Type: types.ResourceTypePeer}, + Destinations: []string{groupIDs[1]}, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) + assert.Contains(t, groups, groupIDs[0]) + assert.Contains(t, groups, groupIDs[1]) + assert.Contains(t, directPeers, peerIDs[3], "direct peer resource should be in directPeers") +} + +func TestCollectGroupChange_RouteLinked(t *testing.T) { + manager, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Create route: PeerGroups=[group0], Groups(distribution)=[group1] + _, err := manager.CreateRoute(ctx, accountID, + netip.MustParsePrefix("10.0.0.0/24"), // network + route.IPv4Network, // type + nil, // domains + "", // peer + []string{groupIDs[0]}, // peerGroups + "test route", // description + "testnet", // netID + false, // masquerade + 9999, // metric + []string{groupIDs[1]}, // groups (distribution) + []string{groupIDs[2]}, // accessControlGroups + true, // enabled + userID, + false, // keepRoute + false, // skipAutoApply + ) + require.NoError(t, err) + + // Changing group0 (peerGroups) should include group0, group1, group2 + groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) + assert.Contains(t, groups, groupIDs[0]) + assert.Contains(t, groups, groupIDs[1]) + assert.Contains(t, groups, groupIDs[2]) + + // Changing group1 (distribution) should include all three + groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) + assert.Contains(t, groups, groupIDs[0]) + assert.Contains(t, groups, groupIDs[1]) + assert.Contains(t, groups, groupIDs[2]) + + // Changing group3 (not in route) should return nothing + groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[3]}) + assert.Empty(t, groups) +} + +func TestCollectGroupChange_RouteWithDirectPeer(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Create route with direct peer + _, err := manager.CreateRoute(ctx, accountID, + netip.MustParsePrefix("10.1.0.0/24"), + route.IPv4Network, + nil, + peerIDs[4], // direct peer + nil, // no peerGroups + "test route peer", + "testnet2", + false, + 9999, + []string{groupIDs[1]}, // distribution groups + nil, + true, + userID, + false, + false, + ) + require.NoError(t, err) + + // Changing group1 should include group1 and direct peer4 + groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) + assert.Contains(t, groups, groupIDs[1]) + assert.Contains(t, directPeers, peerIDs[4]) +} + +func TestCollectGroupChange_NameServerGroupLinked(t *testing.T) { + manager, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Create nameserver group with group0 + _, err := manager.CreateNameServerGroup(ctx, accountID, "ns1", "NS Group 1", + []nbdns.NameServer{{ + IP: netip.MustParseAddr("1.1.1.1"), + NSType: nbdns.UDPNameServerType, + Port: nbdns.DefaultDNSPort, + }}, + []string{groupIDs[0]}, + true, nil, true, userID, false, + ) + require.NoError(t, err) + + // Changing group0 should include group0 (the NS group references it) + groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) + assert.Contains(t, groups, groupIDs[0]) + + // Changing group1 (not in NS group) should not include anything from NS + groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) + assert.Empty(t, groups) +} + +func TestCollectGroupChange_DNSSettingsLinked(t *testing.T) { + manager, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + err := manager.SaveDNSSettings(ctx, accountID, userID, &types.DNSSettings{ + DisabledManagementGroups: []string{groupIDs[2]}, + }) + require.NoError(t, err) + + // Changing group2 should include group2 (from DNS disabled management) + groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[2]}) + assert.Contains(t, groups, groupIDs[2]) + + // Changing group0 (not in DNS settings) should return nothing + groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) + assert.Empty(t, groups) +} + +func TestCollectGroupChange_NetworkRouterLinked(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Create network and router + net1 := &networkTypes.Network{ + ID: "net-test-1", + AccountID: accountID, + Name: "test-network", + } + err := manager.Store.SaveNetwork(ctx, net1) + require.NoError(t, err) + + err = manager.Store.SaveNetworkRouter(ctx, &routerTypes.NetworkRouter{ + ID: "router1", + NetworkID: net1.ID, + AccountID: accountID, + PeerGroups: []string{groupIDs[0]}, + Peer: peerIDs[3], + }) + require.NoError(t, err) + + // Changing group0 should include group0 and direct peer3 + groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) + assert.Contains(t, groups, groupIDs[0]) + assert.Contains(t, directPeers, peerIDs[3]) + + // Changing group1 (not in router) should return nothing + groups, directPeers = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) + assert.Empty(t, groups) + assert.Empty(t, directPeers) +} + +func TestCollectGroupChange_MultipleEntities(t *testing.T) { + manager, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Policy: group0 <-> group1 + _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[0]}, + Destinations: []string{groupIDs[1]}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + // Route: PeerGroups=[group2], distribution=[group3] + _, err = manager.CreateRoute(ctx, accountID, + netip.MustParsePrefix("10.2.0.0/24"), + route.IPv4Network, + nil, + "", + []string{groupIDs[2]}, + "multi route", + "multinet", + false, + 9999, + []string{groupIDs[3]}, + nil, + true, + userID, + false, + false, + ) + require.NoError(t, err) + + // Changing group0 should only pick up policy groups (group0, group1), not route groups + groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) + assert.Contains(t, groups, groupIDs[0]) + assert.Contains(t, groups, groupIDs[1]) + assert.NotContains(t, groups, groupIDs[2], "route groups should not be included for group0 change") + assert.NotContains(t, groups, groupIDs[3], "route groups should not be included for group0 change") + assert.Empty(t, directPeers) + + // Changing group3 should only pick up route groups (group2, group3) + groups, directPeers = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[3]}) + assert.Contains(t, groups, groupIDs[2]) + assert.Contains(t, groups, groupIDs[3]) + assert.NotContains(t, groups, groupIDs[0], "policy groups should not be included for group3 change") + assert.NotContains(t, groups, groupIDs[1], "policy groups should not be included for group3 change") + assert.Empty(t, directPeers) +} + +// ---------- collectPolicyAffectedGroupsAndPeers ---------- + +func TestCollectPolicyAffectedGroups_Basic(t *testing.T) { + policy := &types.Policy{ + Rules: []*types.PolicyRule{ + { + Sources: []string{"g1", "g2"}, + Destinations: []string{"g3"}, + }, + }, + } + groups, directPeers := collectPolicyAffectedGroupsAndPeers(policy) + assert.ElementsMatch(t, []string{"g1", "g2", "g3"}, groups) + assert.Empty(t, directPeers) +} + +func TestCollectPolicyAffectedGroups_WithPeerResources(t *testing.T) { + policy := &types.Policy{ + Rules: []*types.PolicyRule{ + { + Sources: []string{"g1"}, + SourceResource: types.Resource{ID: "p1", Type: types.ResourceTypePeer}, + Destinations: []string{"g2"}, + DestinationResource: types.Resource{ID: "p2", Type: types.ResourceTypePeer}, + }, + }, + } + groups, directPeers := collectPolicyAffectedGroupsAndPeers(policy) + assert.ElementsMatch(t, []string{"g1", "g2"}, groups) + assert.ElementsMatch(t, []string{"p1", "p2"}, directPeers) +} + +func TestCollectPolicyAffectedGroups_NilPolicy(t *testing.T) { + groups, directPeers := collectPolicyAffectedGroupsAndPeers(nil) + assert.Nil(t, groups) + assert.Nil(t, directPeers) +} + +func TestCollectPolicyAffectedGroups_MultipleRules(t *testing.T) { + policy := &types.Policy{ + Rules: []*types.PolicyRule{ + {Sources: []string{"g1"}, Destinations: []string{"g2"}}, + {Sources: []string{"g3"}, Destinations: []string{"g4"}}, + }, + } + groups, _ := collectPolicyAffectedGroupsAndPeers(policy) + assert.ElementsMatch(t, []string{"g1", "g2", "g3", "g4"}, groups) +} + +// ---------- collectRouteAffectedGroupsAndPeers ---------- + +func TestCollectRouteAffectedGroups_Basic(t *testing.T) { + r := &route.Route{ + Groups: []string{"g1"}, + PeerGroups: []string{"g2"}, + AccessControlGroups: []string{"g3"}, + } + groups, directPeers := collectRouteAffectedGroupsAndPeers(r) + assert.ElementsMatch(t, []string{"g1", "g2", "g3"}, groups) + assert.Empty(t, directPeers) +} + +func TestCollectRouteAffectedGroups_WithDirectPeer(t *testing.T) { + r := &route.Route{ + Groups: []string{"g1"}, + Peer: "p1", + } + groups, directPeers := collectRouteAffectedGroupsAndPeers(r) + assert.ElementsMatch(t, []string{"g1"}, groups) + assert.ElementsMatch(t, []string{"p1"}, directPeers) +} + +func TestCollectRouteAffectedGroups_NilRoute(t *testing.T) { + groups, directPeers := collectRouteAffectedGroupsAndPeers(nil) + assert.Nil(t, groups) + assert.Nil(t, directPeers) +} + +// ---------- resolvePeerIDs ---------- + +func TestResolvePeerIDs_GroupsOnly(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // group0 has peer0, group1 has peer1 + result := manager.resolvePeerIDs(ctx, s, accountID, []string{groupIDs[0], groupIDs[1]}, nil) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) +} + +func TestResolvePeerIDs_WithDirectPeers(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // group0 has peer0, plus direct peer2 + result := manager.resolvePeerIDs(ctx, s, accountID, []string{groupIDs[0]}, []string{peerIDs[2]}) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[2]}, result) +} + +func TestResolvePeerIDs_Deduplication(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // peer0 is in group0 and also passed as direct peer -> should appear once + result := manager.resolvePeerIDs(ctx, s, accountID, []string{groupIDs[0]}, []string{peerIDs[0]}) + assert.Len(t, result, 1) + assert.Equal(t, peerIDs[0], result[0]) +} + +func TestResolvePeerIDs_EmptyInputs(t *testing.T) { + manager, s, accountID, _, _ := setupAffectedPeersTest(t) + ctx := context.Background() + + result := manager.resolvePeerIDs(ctx, s, accountID, nil, nil) + assert.Empty(t, result) +} + +// ---------- resolveAffectedPeersForPeerChanges (end-to-end) ---------- + +func TestResolveAffectedPeers_NoPoliciesOrRoutes(t *testing.T) { + manager, s, accountID, peerIDs, _ := setupAffectedPeersTest(t) + ctx := context.Background() + + // No entity references any group, so changing peer0 should yield 0 affected peers + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) + assert.Empty(t, result, "no entities reference any group, should return no affected peers") +} + +func TestResolveAffectedPeers_PolicyBetweenTwoGroups(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Policy: group0 (src) <-> group1 (dst) + _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[0]}, + Destinations: []string{groupIDs[1]}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + // Changing peer0 (in group0) should affect peer0 + peer1 + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) + + // Changing peer1 (in group1) should also affect peer0 + peer1 + result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[1]}) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) + + // Changing peer2 (in group2, not in policy) should return empty + result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[2]}) + assert.Empty(t, result) +} + +func TestResolveAffectedPeers_PolicyThreeGroups(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Policy with multiple sources/destinations: group0,group1 -> group2 + _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[0], groupIDs[1]}, + Destinations: []string{groupIDs[2]}, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + // Changing peer0 (in group0) should affect peer0 + peer1 + peer2 + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1], peerIDs[2]}, result) +} + +func TestResolveAffectedPeers_RoutePeerGroups(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Route: peerGroups=[group0], distribution=[group1] + _, err := manager.CreateRoute(ctx, accountID, + netip.MustParsePrefix("10.3.0.0/24"), + route.IPv4Network, + nil, + "", + []string{groupIDs[0]}, + "test route", + "routenet", + false, + 9999, + []string{groupIDs[1]}, + nil, + true, + userID, + false, + false, + ) + require.NoError(t, err) + + // Changing peer0 (in group0/peerGroups) should affect peer0 + peer1 + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) + + // Changing peer1 (in group1/distribution) should also affect peer0 + peer1 + result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[1]}) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) + + // Changing peer2 (unrelated) should return empty + result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[2]}) + assert.Empty(t, result) +} + +func TestResolveAffectedPeers_RouteWithDirectPeer(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Route with direct peer: peer=peer4, distribution=[group1] + _, err := manager.CreateRoute(ctx, accountID, + netip.MustParsePrefix("10.4.0.0/24"), + route.IPv4Network, + nil, + peerIDs[4], + nil, + "route with peer", + "routenet2", + false, + 9999, + []string{groupIDs[1]}, + nil, + true, + userID, + false, + false, + ) + require.NoError(t, err) + + // Changing peer1 (in distribution group1) should affect peer1 + direct peer4 + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[1]}) + assert.ElementsMatch(t, []string{peerIDs[1], peerIDs[4]}, result) +} + +func TestResolveAffectedPeers_NetworkRouter(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Create network + router with peerGroups=[group0], direct peer=peer3 + net1 := &networkTypes.Network{ + ID: "net-test-2", + AccountID: accountID, + Name: "test-net", + } + err := manager.Store.SaveNetwork(ctx, net1) + require.NoError(t, err) + + err = manager.Store.SaveNetworkRouter(ctx, &routerTypes.NetworkRouter{ + ID: "router-test", + NetworkID: net1.ID, + AccountID: accountID, + PeerGroups: []string{groupIDs[0]}, + Peer: peerIDs[3], + }) + require.NoError(t, err) + + // Changing peer0 (in group0) should affect peer0 + peer3 (direct) + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[3]}, result) +} + +func TestResolveAffectedPeers_NameServerGroup(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // NS group with group0 + _, err := manager.CreateNameServerGroup(ctx, accountID, "ns-test", "NS Test", + []nbdns.NameServer{{ + IP: netip.MustParseAddr("8.8.8.8"), + NSType: nbdns.UDPNameServerType, + Port: nbdns.DefaultDNSPort, + }}, + []string{groupIDs[0]}, + true, nil, true, userID, false, + ) + require.NoError(t, err) + + // Changing peer0 (in group0) should affect peer0 + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) + assert.Contains(t, result, peerIDs[0]) +} + +func TestResolveAffectedPeers_DNSSettings(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // DNS disabled management on group0 + err := manager.SaveDNSSettings(ctx, accountID, userID, &types.DNSSettings{ + DisabledManagementGroups: []string{groupIDs[0]}, + }) + require.NoError(t, err) + + // Changing peer0 (in group0) should affect peer0 + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) + assert.Contains(t, result, peerIDs[0]) +} + +func TestResolveAffectedPeers_PeerInMultipleGroups(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Add peer0 to group1 as well + err := manager.GroupAddPeer(ctx, accountID, groupIDs[1], peerIDs[0]) + require.NoError(t, err) + + // Policy: group0 -> group2 + _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[0]}, + Destinations: []string{groupIDs[2]}, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + // Another policy: group1 -> group3 + _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[1]}, + Destinations: []string{groupIDs[3]}, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + // Changing peer0 (in group0 AND group1) should affect: + // From policy1: group0+group2 -> peer0, peer2 + // From policy2: group1+group3 -> peer0, peer1, peer3 + // Total: peer0, peer1, peer2, peer3 + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1], peerIDs[2], peerIDs[3]}, result) +} + +func TestResolveAffectedPeers_MultipleChangedPeers(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Policy: group0 <-> group1 + _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[0]}, + Destinations: []string{groupIDs[1]}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + // Another policy: group2 <-> group3 + _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[2]}, + Destinations: []string{groupIDs[3]}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + // Changing peer0 AND peer2 at once + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0], peerIDs[2]}) + // peer0 -> policy1 -> peer0, peer1 + // peer2 -> policy2 -> peer2, peer3 + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1], peerIDs[2], peerIDs[3]}, result) +} + +func TestResolveAffectedPeers_SharedGroupAcrossPolicyAndRoute(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Policy: group0 <-> group1 + _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[0]}, + Destinations: []string{groupIDs[1]}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + // Route: distribution=[group0], peerGroups=[group2] + _, err = manager.CreateRoute(ctx, accountID, + netip.MustParsePrefix("10.5.0.0/24"), + route.IPv4Network, + nil, + "", + []string{groupIDs[2]}, + "shared group route", + "sharednet", + false, + 9999, + []string{groupIDs[0]}, + nil, + true, + userID, + false, + false, + ) + require.NoError(t, err) + + // Changing peer0 (in group0) should affect: + // From policy: group0+group1 -> peer0, peer1 + // From route: group0+group2 -> peer0, peer2 + // Total: peer0, peer1, peer2 + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1], peerIDs[2]}, result) +} + +func TestResolveAffectedPeers_EmptyChangedPeers(t *testing.T) { + manager, s, accountID, _, _ := setupAffectedPeersTest(t) + ctx := context.Background() + + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, nil) + assert.Empty(t, result) + + result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{}) + assert.Empty(t, result) +} + +// ---------- Integration: peer changes with full update flow ---------- + +func TestAffectedPeers_GroupUpdateOnlyAffectsLinkedPeers(t *testing.T) { + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + // Delete the default "All <-> All" policy so only our explicit policy matters + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + err := manager.Store.DeletePolicy(ctx, accountID, p.ID) + require.NoError(t, err) + } + + // Create groups + for _, g := range []*types.Group{ + {ID: "ap-grpA", Name: "AP-A", Peers: []string{peer1.ID}}, + {ID: "ap-grpB", Name: "AP-B", Peers: []string{peer2.ID}}, + {ID: "ap-grpC", Name: "AP-C", Peers: []string{peer3.ID}}, + } { + err := manager.CreateGroup(ctx, accountID, userID, g) + require.NoError(t, err) + } + + // Policy: grpA <-> grpB (peer1 <-> peer2). peer3 is NOT in this policy. + _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{"ap-grpA"}, + Destinations: []string{"ap-grpB"}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + // Open update channels for all peers + updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) + updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) + updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + // Verify resolution: changing peer1 should only affect peer1 and peer2 + result := manager.resolveAffectedPeersForPeerChanges(ctx, manager.Store, accountID, []string{peer1.ID}) + assert.ElementsMatch(t, []string{peer1.ID, peer2.ID}, result) + + // Updating grpA to include peer3 should update all 3 peers because after the update + // grpA={peer1,peer3} which is in the policy, plus grpB={peer2} + t.Run("group change updates all peers in policy groups", func(t *testing.T) { + done := make(chan struct{}) + go func() { + peerShouldReceiveUpdate(t, updMsg1) + peerShouldReceiveUpdate(t, updMsg2) + peerShouldReceiveUpdate(t, updMsg3) // peer3 is now in grpA which is in the policy + close(done) + }() + + err := manager.UpdateGroup(ctx, accountID, userID, &types.Group{ + ID: "ap-grpA", + Name: "AP-A", + Peers: []string{peer1.ID, peer3.ID}, // add peer3 to group + }) + assert.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout") + } + }) + + _ = updMsg1 + _ = updMsg2 + _ = updMsg3 +} + +func TestAffectedPeers_UnlinkedGroupChange_NoUpdates(t *testing.T) { + manager, s, accountID, peerIDs, _ := setupAffectedPeersTest(t) + ctx := context.Background() + + // No entities reference any group + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) + assert.Empty(t, result, "unlinked peer change should produce no affected peers") +} + +// ---------- collectPostureCheckAffectedGroupsAndPeers ---------- + +func TestCollectPostureCheckAffected_NoMatch(t *testing.T) { + _, s, accountID, _, _ := setupAffectedPeersTest(t) + ctx := context.Background() + + groups, directPeers := collectPostureCheckAffectedGroupsAndPeers(ctx, s, accountID, "nonexistent-check") + assert.Empty(t, groups) + assert.Empty(t, directPeers) +} + +// ---------- Isolation: unrelated entities don't bleed ---------- + +func TestAffectedPeers_IsolatedPolicies(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Policy A: group0 <-> group1 + _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[0]}, + Destinations: []string{groupIDs[1]}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + // Policy B: group2 <-> group3 + _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[2]}, + Destinations: []string{groupIDs[3]}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + // Changing peer0 should ONLY affect peer0, peer1 (Policy A), NOT peer2, peer3 (Policy B) + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) + assert.NotContains(t, result, peerIDs[2]) + assert.NotContains(t, result, peerIDs[3]) + + // Changing peer2 should ONLY affect peer2, peer3 (Policy B) + result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[2]}) + assert.ElementsMatch(t, []string{peerIDs[2], peerIDs[3]}, result) + assert.NotContains(t, result, peerIDs[0]) + assert.NotContains(t, result, peerIDs[1]) + + // Changing peer4 (not in any policy) should return empty + result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[4]}) + assert.Empty(t, result) +} + +func TestAffectedPeers_IsolatedRouteAndPolicy(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Policy: group0 <-> group1 + _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[0]}, + Destinations: []string{groupIDs[1]}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + // Route: peerGroups=[group2], distribution=[group3] + _, err = manager.CreateRoute(ctx, accountID, + netip.MustParsePrefix("10.6.0.0/24"), + route.IPv4Network, + nil, + "", + []string{groupIDs[2]}, + "isolated route", + "isonet", + false, + 9999, + []string{groupIDs[3]}, + nil, + true, + userID, + false, + false, + ) + require.NoError(t, err) + + // Changing peer0 (policy only) -> peer0, peer1 + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) + assert.NotContains(t, result, peerIDs[2]) + assert.NotContains(t, result, peerIDs[3]) + + // Changing peer2 (route only) -> peer2, peer3 + result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[2]}) + assert.ElementsMatch(t, []string{peerIDs[2], peerIDs[3]}, result) + assert.NotContains(t, result, peerIDs[0]) + assert.NotContains(t, result, peerIDs[1]) +} + +// ---------- Helper: verify no duplicates in resolution ---------- + +func TestResolveAffectedPeers_NoDuplicates(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Add peer0 to multiple groups + err := manager.GroupAddPeer(ctx, accountID, groupIDs[1], peerIDs[0]) + require.NoError(t, err) + err = manager.GroupAddPeer(ctx, accountID, groupIDs[2], peerIDs[0]) + require.NoError(t, err) + + // Policy that references all three groups + _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[0], groupIDs[1]}, + Destinations: []string{groupIDs[2]}, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) + // peer0 is in group0, group1, group2 - should only appear once + count := 0 + for _, id := range result { + if id == peerIDs[0] { + count++ + } + } + assert.Equal(t, 1, count, "peer0 should appear exactly once in results") +} + +// ---------- policyReferencesGroups ---------- + +func TestPolicyReferencesGroups(t *testing.T) { + policy := &types.Policy{ + Rules: []*types.PolicyRule{ + { + Sources: []string{"g1", "g2"}, + Destinations: []string{"g3"}, + }, + }, + } + + tests := []struct { + name string + groupSet map[string]struct{} + want bool + }{ + {"matches source", map[string]struct{}{"g1": {}}, true}, + {"matches destination", map[string]struct{}{"g3": {}}, true}, + {"no match", map[string]struct{}{"g4": {}}, false}, + {"empty set", map[string]struct{}{}, false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := policyReferencesGroups(policy, tt.groupSet) + assert.Equal(t, tt.want, got) + }) + } +} + +func TestRouteReferencesGroups(t *testing.T) { + r := &route.Route{ + Groups: []string{"g1"}, + PeerGroups: []string{"g2"}, + AccessControlGroups: []string{"g3"}, + } + + tests := []struct { + name string + groupSet map[string]struct{} + want bool + }{ + {"matches groups", map[string]struct{}{"g1": {}}, true}, + {"matches peerGroups", map[string]struct{}{"g2": {}}, true}, + {"matches accessControl", map[string]struct{}{"g3": {}}, true}, + {"no match", map[string]struct{}{"g4": {}}, false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := routeReferencesGroups(r, tt.groupSet) + assert.Equal(t, tt.want, got) + }) + } +} + +func TestRouterReferencesGroups(t *testing.T) { + router := &routerTypes.NetworkRouter{ + PeerGroups: []string{"g1", "g2"}, + } + + tests := []struct { + name string + groupSet map[string]struct{} + want bool + }{ + {"matches", map[string]struct{}{"g1": {}}, true}, + {"no match", map[string]struct{}{"g3": {}}, false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := routerReferencesGroups(router, tt.groupSet) + assert.Equal(t, tt.want, got) + }) + } +} + +// ---------- helpers ---------- + +func addPeerToAccount(t *testing.T, manager *DefaultAccountManager, accountID, setupKeyKey string) *nbpeer.Peer { + t.Helper() + + key, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err) + + peer, _, _, err := manager.AddPeer(context.Background(), "", setupKeyKey, "", &nbpeer.Peer{ + Key: key.PublicKey().String(), + Meta: nbpeer.PeerSystemMeta{Hostname: key.PublicKey().String()}, + }, false) + require.NoError(t, err) + return peer +} diff --git a/management/server/posture_checks_test.go b/management/server/posture_checks_test.go index 394f0d896..14bc2c45a 100644 --- a/management/server/posture_checks_test.go +++ b/management/server/posture_checks_test.go @@ -503,21 +503,20 @@ func TestArePostureCheckChangesAffectPeers(t *testing.T) { require.NoError(t, err, "failed to save policy") t.Run("posture check exists and is linked to policy with peers", func(t *testing.T) { - result, err := arePostureCheckChangesAffectPeers(context.Background(), manager.Store, account.Id, postureCheckA.ID) - require.NoError(t, err) - assert.True(t, result) + groupIDs, _ := collectPostureCheckAffectedGroupsAndPeers(context.Background(), manager.Store, account.Id, postureCheckA.ID) + assert.NotEmpty(t, groupIDs) }) t.Run("posture check exists but is not linked to any policy", func(t *testing.T) { - result, err := arePostureCheckChangesAffectPeers(context.Background(), manager.Store, account.Id, postureCheckB.ID) - require.NoError(t, err) - assert.False(t, result) + groupIDs, directPeerIDs := collectPostureCheckAffectedGroupsAndPeers(context.Background(), manager.Store, account.Id, postureCheckB.ID) + assert.Empty(t, groupIDs) + assert.Empty(t, directPeerIDs) }) t.Run("posture check does not exist", func(t *testing.T) { - result, err := arePostureCheckChangesAffectPeers(context.Background(), manager.Store, account.Id, "unknown") - require.NoError(t, err) - assert.False(t, result) + groupIDs, directPeerIDs := collectPostureCheckAffectedGroupsAndPeers(context.Background(), manager.Store, account.Id, "unknown") + assert.Empty(t, groupIDs) + assert.Empty(t, directPeerIDs) }) t.Run("posture check is linked to policy with no peers in source groups", func(t *testing.T) { @@ -526,9 +525,8 @@ func TestArePostureCheckChangesAffectPeers(t *testing.T) { _, err = manager.SavePolicy(context.Background(), account.Id, adminUserID, policy, true) require.NoError(t, err, "failed to update policy") - result, err := arePostureCheckChangesAffectPeers(context.Background(), manager.Store, account.Id, postureCheckA.ID) - require.NoError(t, err) - assert.True(t, result) + groupIDs, _ := collectPostureCheckAffectedGroupsAndPeers(context.Background(), manager.Store, account.Id, postureCheckA.ID) + assert.NotEmpty(t, groupIDs) }) t.Run("posture check is linked to policy with no peers in destination groups", func(t *testing.T) { @@ -537,9 +535,8 @@ func TestArePostureCheckChangesAffectPeers(t *testing.T) { _, err = manager.SavePolicy(context.Background(), account.Id, adminUserID, policy, true) require.NoError(t, err, "failed to update policy") - result, err := arePostureCheckChangesAffectPeers(context.Background(), manager.Store, account.Id, postureCheckA.ID) - require.NoError(t, err) - assert.True(t, result) + groupIDs, _ := collectPostureCheckAffectedGroupsAndPeers(context.Background(), manager.Store, account.Id, postureCheckA.ID) + assert.NotEmpty(t, groupIDs) }) t.Run("posture check is linked to policy but no peers in groups", func(t *testing.T) { @@ -547,9 +544,9 @@ func TestArePostureCheckChangesAffectPeers(t *testing.T) { err = manager.UpdateGroup(context.Background(), account.Id, adminUserID, groupA) require.NoError(t, err, "failed to save groups") - result, err := arePostureCheckChangesAffectPeers(context.Background(), manager.Store, account.Id, postureCheckA.ID) - require.NoError(t, err) - assert.False(t, result) + // The collector returns groups even if they have no peers — the groups are still referenced + groupIDs, _ := collectPostureCheckAffectedGroupsAndPeers(context.Background(), manager.Store, account.Id, postureCheckA.ID) + assert.NotEmpty(t, groupIDs) }) t.Run("posture check is linked to policy with non-existent group", func(t *testing.T) { @@ -558,8 +555,10 @@ func TestArePostureCheckChangesAffectPeers(t *testing.T) { _, err = manager.SavePolicy(context.Background(), account.Id, adminUserID, policy, true) require.NoError(t, err, "failed to update policy") - result, err := arePostureCheckChangesAffectPeers(context.Background(), manager.Store, account.Id, postureCheckA.ID) - require.NoError(t, err) - assert.False(t, result) + // Non-existent groups are filtered out during SavePolicy validation, + // so the saved policy has empty Sources/Destinations + groupIDs, directPeerIDs := collectPostureCheckAffectedGroupsAndPeers(context.Background(), manager.Store, account.Id, postureCheckA.ID) + assert.Empty(t, groupIDs) + assert.Empty(t, directPeerIDs) }) } From 26d778374b5db0bbc889d2205aef1e5ad78c0b84 Mon Sep 17 00:00:00 2001 From: pascal Date: Tue, 28 Apr 2026 16:56:34 +0200 Subject: [PATCH 05/28] clean comments --- .../network_map/controller/controller.go | 8 +- management/server/affected_peers_test.go | 179 ++++-------------- management/server/group.go | 11 +- management/server/peer.go | 11 +- management/server/policy.go | 3 +- management/server/posture_checks.go | 3 +- 6 files changed, 42 insertions(+), 173 deletions(-) diff --git a/management/internals/controllers/network_map/controller/controller.go b/management/internals/controllers/network_map/controller/controller.go index b14d0d81a..8389fcdb8 100644 --- a/management/internals/controllers/network_map/controller/controller.go +++ b/management/internals/controllers/network_map/controller/controller.go @@ -270,8 +270,6 @@ func (c *Controller) UpdateAccountPeers(ctx context.Context, accountID string) e } // UpdateAffectedPeers updates only the specified peers that belong to an account. -// Should be called when a change is known to affect only a subset of peers. -// If peerIDs is empty, this is a no-op. func (c *Controller) UpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error { if len(peerIDs) == 0 { return nil @@ -287,7 +285,6 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s affected[id] = struct{}{} } - // Fast check: any of the affected peers actually connected? hasConnected := false for _, id := range peerIDs { if c.peersUpdateManager.HasChannel(id) { @@ -306,7 +303,6 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s globalStart := time.Now() - // Collect the subset of account peers that are both affected and connected. var peersToUpdate []*nbpeer.Peer for _, peer := range account.Peers { if _, ok := affected[peer.ID]; ok && c.peersUpdateManager.HasChannel(peer.ID) { @@ -504,8 +500,7 @@ func (c *Controller) BufferUpdateAccountPeers(ctx context.Context, accountID str return nil } -// BufferUpdateAffectedPeers accumulates peer IDs across rapid successive calls -// and flushes them in a single sendUpdateForAffectedPeers call after the buffer interval. +// BufferUpdateAffectedPeers accumulates peer IDs and flushes them after the buffer interval. func (c *Controller) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error { if len(peerIDs) == 0 { return nil @@ -518,7 +513,6 @@ func (c *Controller) BufferUpdateAffectedPeers(ctx context.Context, accountID st }) b := bufUpd.(*bufferAffectedUpdate) - // Always accumulate incoming peer IDs (non-blocking). b.addPeerIDs(peerIDs) if !b.sendMu.TryLock() { diff --git a/management/server/affected_peers_test.go b/management/server/affected_peers_test.go index cf0cb033d..bf48b0e4a 100644 --- a/management/server/affected_peers_test.go +++ b/management/server/affected_peers_test.go @@ -20,15 +20,8 @@ import ( "github.com/netbirdio/netbird/route" ) -// setupAffectedPeersTest creates a manager with a clean account and 5 peers, each in its own group. -// Returns the manager, store, account ID, peer IDs, and group IDs. -// Peer layout: -// -// peer0 -> group0 -// peer1 -> group1 -// peer2 -> group2 -// peer3 -> group3 -// peer4 -> group4 +// setupAffectedPeersTest creates a manager with a clean account (default policy deleted) +// and 5 peers, each in its own group: peer0->group0, peer1->group1, ..., peer4->group4. func setupAffectedPeersTest(t *testing.T) (*DefaultAccountManager, store.Store, string, []string, []string) { t.Helper() @@ -41,7 +34,6 @@ func setupAffectedPeersTest(t *testing.T) (*DefaultAccountManager, store.Store, ctx := context.Background() accountID := account.Id - // Delete the default "All <-> All" policy so tests start with a clean slate policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) require.NoError(t, err) for _, p := range policies { @@ -76,16 +68,13 @@ func setupAffectedPeersTest(t *testing.T) (*DefaultAccountManager, store.Store, func affectedGroupID(i int) string { return fmt.Sprintf("affected-grp-%d", i) } func affectedGroupName(i int) string { return fmt.Sprintf("AffectedGroup%d", i) } -// ---------- collectGroupChangeAffectedGroups ---------- - func TestCollectGroupChange_NoEntities(t *testing.T) { _, s, accountID, _, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // group0 is not referenced by any entity groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) - assert.Empty(t, groups, "no entities reference group0, should return no affected groups") - assert.Empty(t, directPeers, "no entities reference group0, should return no direct peers") + assert.Empty(t, groups) + assert.Empty(t, directPeers) } func TestCollectGroupChange_EmptyInput(t *testing.T) { @@ -105,7 +94,6 @@ func TestCollectGroupChange_PolicyLinked(t *testing.T) { manager, s, accountID, _, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Create policy: group0 (src) <-> group1 (dst) _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ @@ -120,17 +108,14 @@ func TestCollectGroupChange_PolicyLinked(t *testing.T) { }, true) require.NoError(t, err) - // Changing group0 should include both group0 and group1 (from the policy rule) groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) assert.Contains(t, groups, groupIDs[0]) assert.Contains(t, groups, groupIDs[1]) - // Changing group1 should also include both groups groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) assert.Contains(t, groups, groupIDs[0]) assert.Contains(t, groups, groupIDs[1]) - // Changing group2 (not in policy) should return nothing groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[2]}) assert.Empty(t, groups) } @@ -139,7 +124,6 @@ func TestCollectGroupChange_PolicyWithDirectPeerResource(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Create policy with direct peer resource _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ @@ -157,46 +141,42 @@ func TestCollectGroupChange_PolicyWithDirectPeerResource(t *testing.T) { groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) assert.Contains(t, groups, groupIDs[0]) assert.Contains(t, groups, groupIDs[1]) - assert.Contains(t, directPeers, peerIDs[3], "direct peer resource should be in directPeers") + assert.Contains(t, directPeers, peerIDs[3]) } func TestCollectGroupChange_RouteLinked(t *testing.T) { manager, s, accountID, _, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Create route: PeerGroups=[group0], Groups(distribution)=[group1] _, err := manager.CreateRoute(ctx, accountID, - netip.MustParsePrefix("10.0.0.0/24"), // network - route.IPv4Network, // type - nil, // domains - "", // peer - []string{groupIDs[0]}, // peerGroups - "test route", // description - "testnet", // netID - false, // masquerade - 9999, // metric - []string{groupIDs[1]}, // groups (distribution) - []string{groupIDs[2]}, // accessControlGroups - true, // enabled + netip.MustParsePrefix("10.0.0.0/24"), + route.IPv4Network, + nil, + "", + []string{groupIDs[0]}, + "test route", + "testnet", + false, + 9999, + []string{groupIDs[1]}, + []string{groupIDs[2]}, + true, userID, - false, // keepRoute - false, // skipAutoApply + false, + false, ) require.NoError(t, err) - // Changing group0 (peerGroups) should include group0, group1, group2 groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) assert.Contains(t, groups, groupIDs[0]) assert.Contains(t, groups, groupIDs[1]) assert.Contains(t, groups, groupIDs[2]) - // Changing group1 (distribution) should include all three groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) assert.Contains(t, groups, groupIDs[0]) assert.Contains(t, groups, groupIDs[1]) assert.Contains(t, groups, groupIDs[2]) - // Changing group3 (not in route) should return nothing groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[3]}) assert.Empty(t, groups) } @@ -205,18 +185,17 @@ func TestCollectGroupChange_RouteWithDirectPeer(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Create route with direct peer _, err := manager.CreateRoute(ctx, accountID, netip.MustParsePrefix("10.1.0.0/24"), route.IPv4Network, nil, - peerIDs[4], // direct peer - nil, // no peerGroups + peerIDs[4], + nil, "test route peer", "testnet2", false, 9999, - []string{groupIDs[1]}, // distribution groups + []string{groupIDs[1]}, nil, true, userID, @@ -225,7 +204,6 @@ func TestCollectGroupChange_RouteWithDirectPeer(t *testing.T) { ) require.NoError(t, err) - // Changing group1 should include group1 and direct peer4 groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) assert.Contains(t, groups, groupIDs[1]) assert.Contains(t, directPeers, peerIDs[4]) @@ -235,7 +213,6 @@ func TestCollectGroupChange_NameServerGroupLinked(t *testing.T) { manager, s, accountID, _, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Create nameserver group with group0 _, err := manager.CreateNameServerGroup(ctx, accountID, "ns1", "NS Group 1", []nbdns.NameServer{{ IP: netip.MustParseAddr("1.1.1.1"), @@ -247,11 +224,9 @@ func TestCollectGroupChange_NameServerGroupLinked(t *testing.T) { ) require.NoError(t, err) - // Changing group0 should include group0 (the NS group references it) groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) assert.Contains(t, groups, groupIDs[0]) - // Changing group1 (not in NS group) should not include anything from NS groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) assert.Empty(t, groups) } @@ -265,11 +240,9 @@ func TestCollectGroupChange_DNSSettingsLinked(t *testing.T) { }) require.NoError(t, err) - // Changing group2 should include group2 (from DNS disabled management) groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[2]}) assert.Contains(t, groups, groupIDs[2]) - // Changing group0 (not in DNS settings) should return nothing groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) assert.Empty(t, groups) } @@ -278,7 +251,6 @@ func TestCollectGroupChange_NetworkRouterLinked(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Create network and router net1 := &networkTypes.Network{ ID: "net-test-1", AccountID: accountID, @@ -296,12 +268,10 @@ func TestCollectGroupChange_NetworkRouterLinked(t *testing.T) { }) require.NoError(t, err) - // Changing group0 should include group0 and direct peer3 groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) assert.Contains(t, groups, groupIDs[0]) assert.Contains(t, directPeers, peerIDs[3]) - // Changing group1 (not in router) should return nothing groups, directPeers = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) assert.Empty(t, groups) assert.Empty(t, directPeers) @@ -311,7 +281,6 @@ func TestCollectGroupChange_MultipleEntities(t *testing.T) { manager, s, accountID, _, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Policy: group0 <-> group1 _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ @@ -326,7 +295,6 @@ func TestCollectGroupChange_MultipleEntities(t *testing.T) { }, true) require.NoError(t, err) - // Route: PeerGroups=[group2], distribution=[group3] _, err = manager.CreateRoute(ctx, accountID, netip.MustParsePrefix("10.2.0.0/24"), route.IPv4Network, @@ -346,25 +314,21 @@ func TestCollectGroupChange_MultipleEntities(t *testing.T) { ) require.NoError(t, err) - // Changing group0 should only pick up policy groups (group0, group1), not route groups groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) assert.Contains(t, groups, groupIDs[0]) assert.Contains(t, groups, groupIDs[1]) - assert.NotContains(t, groups, groupIDs[2], "route groups should not be included for group0 change") - assert.NotContains(t, groups, groupIDs[3], "route groups should not be included for group0 change") + assert.NotContains(t, groups, groupIDs[2]) + assert.NotContains(t, groups, groupIDs[3]) assert.Empty(t, directPeers) - // Changing group3 should only pick up route groups (group2, group3) groups, directPeers = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[3]}) assert.Contains(t, groups, groupIDs[2]) assert.Contains(t, groups, groupIDs[3]) - assert.NotContains(t, groups, groupIDs[0], "policy groups should not be included for group3 change") - assert.NotContains(t, groups, groupIDs[1], "policy groups should not be included for group3 change") + assert.NotContains(t, groups, groupIDs[0]) + assert.NotContains(t, groups, groupIDs[1]) assert.Empty(t, directPeers) } -// ---------- collectPolicyAffectedGroupsAndPeers ---------- - func TestCollectPolicyAffectedGroups_Basic(t *testing.T) { policy := &types.Policy{ Rules: []*types.PolicyRule{ @@ -412,8 +376,6 @@ func TestCollectPolicyAffectedGroups_MultipleRules(t *testing.T) { assert.ElementsMatch(t, []string{"g1", "g2", "g3", "g4"}, groups) } -// ---------- collectRouteAffectedGroupsAndPeers ---------- - func TestCollectRouteAffectedGroups_Basic(t *testing.T) { r := &route.Route{ Groups: []string{"g1"}, @@ -441,13 +403,10 @@ func TestCollectRouteAffectedGroups_NilRoute(t *testing.T) { assert.Nil(t, directPeers) } -// ---------- resolvePeerIDs ---------- - func TestResolvePeerIDs_GroupsOnly(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // group0 has peer0, group1 has peer1 result := manager.resolvePeerIDs(ctx, s, accountID, []string{groupIDs[0], groupIDs[1]}, nil) assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) } @@ -456,7 +415,6 @@ func TestResolvePeerIDs_WithDirectPeers(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // group0 has peer0, plus direct peer2 result := manager.resolvePeerIDs(ctx, s, accountID, []string{groupIDs[0]}, []string{peerIDs[2]}) assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[2]}, result) } @@ -465,7 +423,6 @@ func TestResolvePeerIDs_Deduplication(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // peer0 is in group0 and also passed as direct peer -> should appear once result := manager.resolvePeerIDs(ctx, s, accountID, []string{groupIDs[0]}, []string{peerIDs[0]}) assert.Len(t, result, 1) assert.Equal(t, peerIDs[0], result[0]) @@ -479,22 +436,18 @@ func TestResolvePeerIDs_EmptyInputs(t *testing.T) { assert.Empty(t, result) } -// ---------- resolveAffectedPeersForPeerChanges (end-to-end) ---------- - func TestResolveAffectedPeers_NoPoliciesOrRoutes(t *testing.T) { manager, s, accountID, peerIDs, _ := setupAffectedPeersTest(t) ctx := context.Background() - // No entity references any group, so changing peer0 should yield 0 affected peers result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) - assert.Empty(t, result, "no entities reference any group, should return no affected peers") + assert.Empty(t, result) } func TestResolveAffectedPeers_PolicyBetweenTwoGroups(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Policy: group0 (src) <-> group1 (dst) _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ @@ -509,15 +462,12 @@ func TestResolveAffectedPeers_PolicyBetweenTwoGroups(t *testing.T) { }, true) require.NoError(t, err) - // Changing peer0 (in group0) should affect peer0 + peer1 result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) - // Changing peer1 (in group1) should also affect peer0 + peer1 result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[1]}) assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) - // Changing peer2 (in group2, not in policy) should return empty result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[2]}) assert.Empty(t, result) } @@ -526,7 +476,6 @@ func TestResolveAffectedPeers_PolicyThreeGroups(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Policy with multiple sources/destinations: group0,group1 -> group2 _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ @@ -540,7 +489,6 @@ func TestResolveAffectedPeers_PolicyThreeGroups(t *testing.T) { }, true) require.NoError(t, err) - // Changing peer0 (in group0) should affect peer0 + peer1 + peer2 result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1], peerIDs[2]}, result) } @@ -549,7 +497,6 @@ func TestResolveAffectedPeers_RoutePeerGroups(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Route: peerGroups=[group0], distribution=[group1] _, err := manager.CreateRoute(ctx, accountID, netip.MustParsePrefix("10.3.0.0/24"), route.IPv4Network, @@ -569,15 +516,12 @@ func TestResolveAffectedPeers_RoutePeerGroups(t *testing.T) { ) require.NoError(t, err) - // Changing peer0 (in group0/peerGroups) should affect peer0 + peer1 result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) - // Changing peer1 (in group1/distribution) should also affect peer0 + peer1 result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[1]}) assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) - // Changing peer2 (unrelated) should return empty result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[2]}) assert.Empty(t, result) } @@ -586,7 +530,6 @@ func TestResolveAffectedPeers_RouteWithDirectPeer(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Route with direct peer: peer=peer4, distribution=[group1] _, err := manager.CreateRoute(ctx, accountID, netip.MustParsePrefix("10.4.0.0/24"), route.IPv4Network, @@ -606,7 +549,6 @@ func TestResolveAffectedPeers_RouteWithDirectPeer(t *testing.T) { ) require.NoError(t, err) - // Changing peer1 (in distribution group1) should affect peer1 + direct peer4 result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[1]}) assert.ElementsMatch(t, []string{peerIDs[1], peerIDs[4]}, result) } @@ -615,7 +557,6 @@ func TestResolveAffectedPeers_NetworkRouter(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Create network + router with peerGroups=[group0], direct peer=peer3 net1 := &networkTypes.Network{ ID: "net-test-2", AccountID: accountID, @@ -633,7 +574,6 @@ func TestResolveAffectedPeers_NetworkRouter(t *testing.T) { }) require.NoError(t, err) - // Changing peer0 (in group0) should affect peer0 + peer3 (direct) result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[3]}, result) } @@ -642,7 +582,6 @@ func TestResolveAffectedPeers_NameServerGroup(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // NS group with group0 _, err := manager.CreateNameServerGroup(ctx, accountID, "ns-test", "NS Test", []nbdns.NameServer{{ IP: netip.MustParseAddr("8.8.8.8"), @@ -654,7 +593,6 @@ func TestResolveAffectedPeers_NameServerGroup(t *testing.T) { ) require.NoError(t, err) - // Changing peer0 (in group0) should affect peer0 result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) assert.Contains(t, result, peerIDs[0]) } @@ -663,13 +601,11 @@ func TestResolveAffectedPeers_DNSSettings(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // DNS disabled management on group0 err := manager.SaveDNSSettings(ctx, accountID, userID, &types.DNSSettings{ DisabledManagementGroups: []string{groupIDs[0]}, }) require.NoError(t, err) - // Changing peer0 (in group0) should affect peer0 result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) assert.Contains(t, result, peerIDs[0]) } @@ -678,11 +614,9 @@ func TestResolveAffectedPeers_PeerInMultipleGroups(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Add peer0 to group1 as well err := manager.GroupAddPeer(ctx, accountID, groupIDs[1], peerIDs[0]) require.NoError(t, err) - // Policy: group0 -> group2 _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ @@ -696,7 +630,6 @@ func TestResolveAffectedPeers_PeerInMultipleGroups(t *testing.T) { }, true) require.NoError(t, err) - // Another policy: group1 -> group3 _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ @@ -710,10 +643,7 @@ func TestResolveAffectedPeers_PeerInMultipleGroups(t *testing.T) { }, true) require.NoError(t, err) - // Changing peer0 (in group0 AND group1) should affect: - // From policy1: group0+group2 -> peer0, peer2 - // From policy2: group1+group3 -> peer0, peer1, peer3 - // Total: peer0, peer1, peer2, peer3 + // peer0 is in group0 AND group1, so both policies apply result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1], peerIDs[2], peerIDs[3]}, result) } @@ -722,7 +652,6 @@ func TestResolveAffectedPeers_MultipleChangedPeers(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Policy: group0 <-> group1 _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ @@ -737,7 +666,6 @@ func TestResolveAffectedPeers_MultipleChangedPeers(t *testing.T) { }, true) require.NoError(t, err) - // Another policy: group2 <-> group3 _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ @@ -752,10 +680,7 @@ func TestResolveAffectedPeers_MultipleChangedPeers(t *testing.T) { }, true) require.NoError(t, err) - // Changing peer0 AND peer2 at once result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0], peerIDs[2]}) - // peer0 -> policy1 -> peer0, peer1 - // peer2 -> policy2 -> peer2, peer3 assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1], peerIDs[2], peerIDs[3]}, result) } @@ -763,7 +688,6 @@ func TestResolveAffectedPeers_SharedGroupAcrossPolicyAndRoute(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Policy: group0 <-> group1 _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ @@ -778,7 +702,6 @@ func TestResolveAffectedPeers_SharedGroupAcrossPolicyAndRoute(t *testing.T) { }, true) require.NoError(t, err) - // Route: distribution=[group0], peerGroups=[group2] _, err = manager.CreateRoute(ctx, accountID, netip.MustParsePrefix("10.5.0.0/24"), route.IPv4Network, @@ -798,10 +721,7 @@ func TestResolveAffectedPeers_SharedGroupAcrossPolicyAndRoute(t *testing.T) { ) require.NoError(t, err) - // Changing peer0 (in group0) should affect: - // From policy: group0+group1 -> peer0, peer1 - // From route: group0+group2 -> peer0, peer2 - // Total: peer0, peer1, peer2 + // group0 is shared: policy gives peer0+peer1, route gives peer0+peer2 result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1], peerIDs[2]}, result) } @@ -817,14 +737,11 @@ func TestResolveAffectedPeers_EmptyChangedPeers(t *testing.T) { assert.Empty(t, result) } -// ---------- Integration: peer changes with full update flow ---------- - func TestAffectedPeers_GroupUpdateOnlyAffectsLinkedPeers(t *testing.T) { manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) ctx := context.Background() accountID := account.Id - // Delete the default "All <-> All" policy so only our explicit policy matters policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) require.NoError(t, err) for _, p := range policies { @@ -832,7 +749,6 @@ func TestAffectedPeers_GroupUpdateOnlyAffectsLinkedPeers(t *testing.T) { require.NoError(t, err) } - // Create groups for _, g := range []*types.Group{ {ID: "ap-grpA", Name: "AP-A", Peers: []string{peer1.ID}}, {ID: "ap-grpB", Name: "AP-B", Peers: []string{peer2.ID}}, @@ -842,7 +758,6 @@ func TestAffectedPeers_GroupUpdateOnlyAffectsLinkedPeers(t *testing.T) { require.NoError(t, err) } - // Policy: grpA <-> grpB (peer1 <-> peer2). peer3 is NOT in this policy. _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ @@ -857,7 +772,6 @@ func TestAffectedPeers_GroupUpdateOnlyAffectsLinkedPeers(t *testing.T) { }, true) require.NoError(t, err) - // Open update channels for all peers updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) @@ -867,25 +781,23 @@ func TestAffectedPeers_GroupUpdateOnlyAffectsLinkedPeers(t *testing.T) { updateManager.CloseChannel(ctx, peer3.ID) }) - // Verify resolution: changing peer1 should only affect peer1 and peer2 result := manager.resolveAffectedPeersForPeerChanges(ctx, manager.Store, accountID, []string{peer1.ID}) assert.ElementsMatch(t, []string{peer1.ID, peer2.ID}, result) - // Updating grpA to include peer3 should update all 3 peers because after the update - // grpA={peer1,peer3} which is in the policy, plus grpB={peer2} + // Adding peer3 to grpA makes it part of the policy, so all 3 peers get updated t.Run("group change updates all peers in policy groups", func(t *testing.T) { done := make(chan struct{}) go func() { peerShouldReceiveUpdate(t, updMsg1) peerShouldReceiveUpdate(t, updMsg2) - peerShouldReceiveUpdate(t, updMsg3) // peer3 is now in grpA which is in the policy + peerShouldReceiveUpdate(t, updMsg3) close(done) }() err := manager.UpdateGroup(ctx, accountID, userID, &types.Group{ ID: "ap-grpA", Name: "AP-A", - Peers: []string{peer1.ID, peer3.ID}, // add peer3 to group + Peers: []string{peer1.ID, peer3.ID}, }) assert.NoError(t, err) @@ -905,13 +817,10 @@ func TestAffectedPeers_UnlinkedGroupChange_NoUpdates(t *testing.T) { manager, s, accountID, peerIDs, _ := setupAffectedPeersTest(t) ctx := context.Background() - // No entities reference any group result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) - assert.Empty(t, result, "unlinked peer change should produce no affected peers") + assert.Empty(t, result) } -// ---------- collectPostureCheckAffectedGroupsAndPeers ---------- - func TestCollectPostureCheckAffected_NoMatch(t *testing.T) { _, s, accountID, _, _ := setupAffectedPeersTest(t) ctx := context.Background() @@ -921,13 +830,10 @@ func TestCollectPostureCheckAffected_NoMatch(t *testing.T) { assert.Empty(t, directPeers) } -// ---------- Isolation: unrelated entities don't bleed ---------- - func TestAffectedPeers_IsolatedPolicies(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Policy A: group0 <-> group1 _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ @@ -942,7 +848,6 @@ func TestAffectedPeers_IsolatedPolicies(t *testing.T) { }, true) require.NoError(t, err) - // Policy B: group2 <-> group3 _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ @@ -957,19 +862,16 @@ func TestAffectedPeers_IsolatedPolicies(t *testing.T) { }, true) require.NoError(t, err) - // Changing peer0 should ONLY affect peer0, peer1 (Policy A), NOT peer2, peer3 (Policy B) result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) assert.NotContains(t, result, peerIDs[2]) assert.NotContains(t, result, peerIDs[3]) - // Changing peer2 should ONLY affect peer2, peer3 (Policy B) result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[2]}) assert.ElementsMatch(t, []string{peerIDs[2], peerIDs[3]}, result) assert.NotContains(t, result, peerIDs[0]) assert.NotContains(t, result, peerIDs[1]) - // Changing peer4 (not in any policy) should return empty result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[4]}) assert.Empty(t, result) } @@ -978,7 +880,6 @@ func TestAffectedPeers_IsolatedRouteAndPolicy(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Policy: group0 <-> group1 _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ @@ -993,7 +894,6 @@ func TestAffectedPeers_IsolatedRouteAndPolicy(t *testing.T) { }, true) require.NoError(t, err) - // Route: peerGroups=[group2], distribution=[group3] _, err = manager.CreateRoute(ctx, accountID, netip.MustParsePrefix("10.6.0.0/24"), route.IPv4Network, @@ -1013,32 +913,26 @@ func TestAffectedPeers_IsolatedRouteAndPolicy(t *testing.T) { ) require.NoError(t, err) - // Changing peer0 (policy only) -> peer0, peer1 result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result) assert.NotContains(t, result, peerIDs[2]) assert.NotContains(t, result, peerIDs[3]) - // Changing peer2 (route only) -> peer2, peer3 result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[2]}) assert.ElementsMatch(t, []string{peerIDs[2], peerIDs[3]}, result) assert.NotContains(t, result, peerIDs[0]) assert.NotContains(t, result, peerIDs[1]) } -// ---------- Helper: verify no duplicates in resolution ---------- - func TestResolveAffectedPeers_NoDuplicates(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - // Add peer0 to multiple groups err := manager.GroupAddPeer(ctx, accountID, groupIDs[1], peerIDs[0]) require.NoError(t, err) err = manager.GroupAddPeer(ctx, accountID, groupIDs[2], peerIDs[0]) require.NoError(t, err) - // Policy that references all three groups _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ @@ -1053,18 +947,15 @@ func TestResolveAffectedPeers_NoDuplicates(t *testing.T) { require.NoError(t, err) result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) - // peer0 is in group0, group1, group2 - should only appear once count := 0 for _, id := range result { if id == peerIDs[0] { count++ } } - assert.Equal(t, 1, count, "peer0 should appear exactly once in results") + assert.Equal(t, 1, count, "peer0 should appear exactly once") } -// ---------- policyReferencesGroups ---------- - func TestPolicyReferencesGroups(t *testing.T) { policy := &types.Policy{ Rules: []*types.PolicyRule{ @@ -1142,8 +1033,6 @@ func TestRouterReferencesGroups(t *testing.T) { } } -// ---------- helpers ---------- - func addPeerToAccount(t *testing.T, manager *DefaultAccountManager, accountID, setupKeyKey string) *nbpeer.Peer { t.Helper() diff --git a/management/server/group.go b/management/server/group.go index cc19cf1a4..cf22e5221 100644 --- a/management/server/group.go +++ b/management/server/group.go @@ -821,10 +821,8 @@ func areGroupChangesAffectPeers(ctx context.Context, transaction store.Store, ac return false, nil } -// collectGroupChangeAffectedGroups walks all entities that reference the changed groups -// and collects the full set of affected group IDs and direct peer IDs. -// This ensures that when a group changes, we update not just the peers in that group -// but also peers in other groups that share policies, routes, DNS, or nameserver configs. +// collectGroupChangeAffectedGroups walks policies, routes, nameservers, DNS settings, +// and network routers to collect all group IDs and direct peer IDs affected by the changed groups. func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedGroupIDs []string) (allGroupIDs []string, directPeerIDs []string) { if len(changedGroupIDs) == 0 { return nil, nil @@ -841,7 +839,6 @@ func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Sto peerSet := make(map[string]struct{}) - // Policies: collect all rule groups + direct peer resources from policies that reference any changed group policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) if err != nil { log.WithContext(ctx).Errorf("failed to get policies for group change resolution: %v", err) @@ -867,7 +864,6 @@ func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Sto } } - // Routes: collect all groups + direct peer from routes that reference any changed group routes, err := transaction.GetAccountRoutes(ctx, store.LockingStrengthNone, accountID) if err != nil { log.WithContext(ctx).Errorf("failed to get routes for group change resolution: %v", err) @@ -893,7 +889,6 @@ func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Sto } } - // Nameserver groups: collect groups from NS groups that reference any changed group nsGroups, err := transaction.GetAccountNameServerGroups(ctx, store.LockingStrengthNone, accountID) if err != nil { log.WithContext(ctx).Errorf("failed to get nameserver groups for group change resolution: %v", err) @@ -911,7 +906,6 @@ func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Sto } } - // DNS settings: if any changed group is in DisabledManagementGroups, include those groups dnsSettings, err := transaction.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID) if err != nil { log.WithContext(ctx).Errorf("failed to get DNS settings for group change resolution: %v", err) @@ -924,7 +918,6 @@ func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Sto } } - // Network routers: collect peer groups + direct peer from routers that reference any changed group routers, err := transaction.GetNetworkRoutersByAccountID(ctx, store.LockingStrengthNone, accountID) if err != nil { log.WithContext(ctx).Errorf("failed to get network routers for group change resolution: %v", err) diff --git a/management/server/peer.go b/management/server/peer.go index ee3a6369f..0d7689619 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -1308,14 +1308,12 @@ func (am *DefaultAccountManager) UpdateAccountPeers(ctx context.Context, account } // UpdateAffectedPeers updates only the specified peers that belong to an account. -// Should be called when a change is known to affect only a subset of peers. func (am *DefaultAccountManager) UpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) { log.WithContext(ctx).Tracef("UpdateAffectedPeers: %d peers for account %s", len(peerIDs), accountID) _ = am.networkMapController.UpdateAffectedPeers(ctx, accountID, peerIDs) } -// resolvePeerIDs resolves a set of group IDs and direct peer IDs into a -// deduplicated list of peer IDs suitable for UpdateAffectedPeers. +// resolvePeerIDs resolves group IDs and direct peer IDs into a deduplicated peer ID list. func (am *DefaultAccountManager) resolvePeerIDs(ctx context.Context, s store.Store, accountID string, groupIDs []string, directPeerIDs []string) []string { peerIDs, err := s.GetPeerIDsByGroups(ctx, accountID, groupIDs) if err != nil { @@ -1343,15 +1341,12 @@ func (am *DefaultAccountManager) resolvePeerIDs(ctx context.Context, s store.Sto return peerIDs } -// BufferUpdateAffectedPeers accumulates peer IDs across rapid successive calls -// and flushes them in a single update after the buffer interval. +// BufferUpdateAffectedPeers accumulates peer IDs and flushes them after the buffer interval. func (am *DefaultAccountManager) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) { _ = am.networkMapController.BufferUpdateAffectedPeers(ctx, accountID, peerIDs) } -// resolveAffectedPeersForPeerChanges resolves changed peer IDs into the full set of -// affected peers: finds groups containing the changed peers, walks all entity linkages, -// and resolves back to peer IDs. +// resolveAffectedPeersForPeerChanges resolves changed peer IDs into the full set of affected peer IDs. func (am *DefaultAccountManager) resolveAffectedPeersForPeerChanges(ctx context.Context, s store.Store, accountID string, changedPeerIDs []string) []string { groupIDs, err := s.GetGroupIDsByPeerIDs(ctx, accountID, changedPeerIDs) if err != nil { diff --git a/management/server/policy.go b/management/server/policy.go index 9ba4f98dd..0404a9208 100644 --- a/management/server/policy.go +++ b/management/server/policy.go @@ -150,8 +150,7 @@ func (am *DefaultAccountManager) ListPolicies(ctx context.Context, accountID, us return am.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) } -// collectPolicyAffectedGroupsAndPeers returns the group IDs and direct peer IDs -// referenced by the given policies' rules. +// collectPolicyAffectedGroupsAndPeers returns group IDs and direct peer IDs from the given policies. func collectPolicyAffectedGroupsAndPeers(policies ...*types.Policy) (groupIDs []string, directPeerIDs []string) { for _, policy := range policies { if policy == nil { diff --git a/management/server/posture_checks.go b/management/server/posture_checks.go index bbf4ed198..cc6abc943 100644 --- a/management/server/posture_checks.go +++ b/management/server/posture_checks.go @@ -130,8 +130,7 @@ func (am *DefaultAccountManager) ListPostureChecks(ctx context.Context, accountI return am.Store.GetAccountPostureChecks(ctx, store.LockingStrengthNone, accountID) } -// collectPostureCheckAffectedGroupsAndPeers finds all policies referencing the given posture check -// and collects their affected group IDs and direct peer IDs. +// collectPostureCheckAffectedGroupsAndPeers returns group IDs and peer IDs from policies referencing the posture check. func collectPostureCheckAffectedGroupsAndPeers(ctx context.Context, transaction store.Store, accountID, postureCheckID string) (groupIDs []string, directPeerIDs []string) { policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) if err != nil { From 0bfccd65d2c688042854357cb3aa881187bf43e7 Mon Sep 17 00:00:00 2001 From: pascal Date: Thu, 30 Apr 2026 16:20:41 +0200 Subject: [PATCH 06/28] add to networks modules --- management/server/networks/manager.go | 121 +++++++++++- .../server/networks/resources/manager.go | 169 ++++++++++++++++- management/server/networks/routers/manager.go | 178 +++++++++++++++++- 3 files changed, 459 insertions(+), 9 deletions(-) diff --git a/management/server/networks/manager.go b/management/server/networks/manager.go index b6706ca45..17ea0ddaa 100644 --- a/management/server/networks/manager.go +++ b/management/server/networks/manager.go @@ -5,6 +5,7 @@ import ( "fmt" "github.com/rs/xid" + log "github.com/sirupsen/logrus" "github.com/netbirdio/netbird/management/server/account" "github.com/netbirdio/netbird/management/server/activity" @@ -15,6 +16,7 @@ import ( "github.com/netbirdio/netbird/management/server/permissions/modules" "github.com/netbirdio/netbird/management/server/permissions/operations" "github.com/netbirdio/netbird/management/server/store" + nbTypes "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/shared/management/status" ) @@ -111,6 +113,14 @@ func (m *managerImpl) UpdateNetwork(ctx context.Context, userID string, network return network, m.store.SaveNetwork(ctx, network) } +// networkAffectedPeersData holds data loaded inside the transaction for affected peer resolution. +type networkAffectedPeersData struct { + resourceGroupIDs []string + routerPeerGroups []string + directPeerIDs []string + policies []*nbTypes.Policy +} + func (m *managerImpl) DeleteNetwork(ctx context.Context, accountID, userID, networkID string) error { ok, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.Networks, operations.Delete) if err != nil { @@ -126,13 +136,22 @@ func (m *managerImpl) DeleteNetwork(ctx context.Context, accountID, userID, netw } var eventsToStore []func() + var affectedData *networkAffectedPeersData err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { resources, err := transaction.GetNetworkResourcesByNetID(ctx, store.LockingStrengthUpdate, accountID, networkID) if err != nil { return fmt.Errorf("failed to get resources in network: %w", err) } + var resourceGroupIDs []string for _, resource := range resources { + groups, err := transaction.GetResourceGroups(ctx, store.LockingStrengthNone, accountID, resource.ID) + if err == nil { + for _, g := range groups { + resourceGroupIDs = append(resourceGroupIDs, g.ID) + } + } + event, err := m.resourcesManager.DeleteResourceInTransaction(ctx, transaction, accountID, userID, networkID, resource.ID) if err != nil { return fmt.Errorf("failed to delete resource: %w", err) @@ -140,12 +159,19 @@ func (m *managerImpl) DeleteNetwork(ctx context.Context, accountID, userID, netw eventsToStore = append(eventsToStore, event...) } - routers, err := transaction.GetNetworkRoutersByNetID(ctx, store.LockingStrengthUpdate, accountID, networkID) + netRouters, err := transaction.GetNetworkRoutersByNetID(ctx, store.LockingStrengthUpdate, accountID, networkID) if err != nil { return fmt.Errorf("failed to get routers in network: %w", err) } - for _, router := range routers { + var routerPeerGroups []string + var directPeerIDs []string + for _, router := range netRouters { + routerPeerGroups = append(routerPeerGroups, router.PeerGroups...) + if router.Peer != "" { + directPeerIDs = append(directPeerIDs, router.Peer) + } + event, err := m.routersManager.DeleteRouterInTransaction(ctx, transaction, accountID, userID, networkID, router.ID) if err != nil { return fmt.Errorf("failed to delete router: %w", err) @@ -153,6 +179,24 @@ func (m *managerImpl) DeleteNetwork(ctx context.Context, accountID, userID, netw eventsToStore = append(eventsToStore, event) } + // load policies before deleting so group memberships are still present + var policies []*nbTypes.Policy + if len(resourceGroupIDs) > 0 { + policies, err = transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get policies for affected peers: %v", err) + } + } + + if len(resourceGroupIDs) > 0 || len(routerPeerGroups) > 0 || len(directPeerIDs) > 0 { + affectedData = &networkAffectedPeersData{ + resourceGroupIDs: resourceGroupIDs, + routerPeerGroups: routerPeerGroups, + directPeerIDs: directPeerIDs, + policies: policies, + } + } + err = transaction.DeleteNetwork(ctx, accountID, networkID) if err != nil { return fmt.Errorf("failed to delete network: %w", err) @@ -177,11 +221,82 @@ func (m *managerImpl) DeleteNetwork(ctx context.Context, accountID, userID, netw event() } - go m.accountManager.UpdateAccountPeers(ctx, accountID) + if affectedData != nil { + affectedPeerIDs := resolveNetworkAffectedPeers(ctx, m.store, accountID, affectedData) + if len(affectedPeerIDs) > 0 { + go m.accountManager.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } + } return nil } +// resolveNetworkAffectedPeers computes affected peer IDs from preloaded data outside the transaction. +func resolveNetworkAffectedPeers(ctx context.Context, s store.Store, accountID string, data *networkAffectedPeersData) []string { + groupSet := make(map[string]struct{}) + + for _, gID := range data.routerPeerGroups { + groupSet[gID] = struct{}{} + } + + if len(data.resourceGroupIDs) > 0 { + destSet := make(map[string]struct{}, len(data.resourceGroupIDs)) + for _, gID := range data.resourceGroupIDs { + destSet[gID] = struct{}{} + groupSet[gID] = struct{}{} + } + + for _, policy := range data.policies { + if policy == nil || !policy.Enabled { + continue + } + for _, rule := range policy.Rules { + if rule == nil || !rule.Enabled { + continue + } + for _, gID := range rule.Destinations { + if _, ok := destSet[gID]; ok { + for _, srcGID := range rule.Sources { + groupSet[srcGID] = struct{}{} + } + break + } + } + } + } + } + + if len(groupSet) == 0 && len(data.directPeerIDs) == 0 { + return nil + } + + groupIDs := make([]string, 0, len(groupSet)) + for gID := range groupSet { + groupIDs = append(groupIDs, gID) + } + + peerIDs, err := s.GetPeerIDsByGroups(ctx, accountID, groupIDs) + if err != nil { + log.WithContext(ctx).Errorf("failed to resolve peer IDs: %v", err) + return nil + } + + if len(data.directPeerIDs) > 0 { + seen := make(map[string]struct{}, len(peerIDs)) + for _, id := range peerIDs { + seen[id] = struct{}{} + } + for _, id := range data.directPeerIDs { + if _, exists := seen[id]; !exists { + peerIDs = append(peerIDs, id) + seen[id] = struct{}{} + } + } + } + + return peerIDs +} + func NewManagerMock() Manager { return &mockManager{} } diff --git a/management/server/networks/resources/manager.go b/management/server/networks/resources/manager.go index 86f9b6579..890f417d5 100644 --- a/management/server/networks/resources/manager.go +++ b/management/server/networks/resources/manager.go @@ -114,6 +114,7 @@ func (m *managerImpl) CreateResource(ctx context.Context, userID string, resourc } var eventsToStore []func() + var affectedData *resourceAffectedPeersData err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { _, err = transaction.GetNetworkResourceByName(ctx, store.LockingStrengthNone, resource.AccountID, resource.Name) if err == nil { @@ -152,6 +153,11 @@ func (m *managerImpl) CreateResource(ctx context.Context, userID string, resourc return fmt.Errorf("failed to increment network serial: %w", err) } + affectedData, err = loadResourceAffectedPeersData(ctx, transaction, resource.AccountID, resource.NetworkID, resource.GroupIDs) + if err != nil { + log.WithContext(ctx).Errorf("failed to load affected peers data: %v", err) + } + return nil }) if err != nil { @@ -162,7 +168,9 @@ func (m *managerImpl) CreateResource(ctx context.Context, userID string, resourc event() } - go m.accountManager.UpdateAccountPeers(ctx, resource.AccountID) + if affectedPeerIDs := m.resolveResourceAffectedPeers(ctx, resource.AccountID, affectedData); len(affectedPeerIDs) > 0 { + go m.accountManager.UpdateAffectedPeers(ctx, resource.AccountID, affectedPeerIDs) + } return resource, nil } @@ -207,6 +215,7 @@ func (m *managerImpl) UpdateResource(ctx context.Context, userID string, resourc resource.Prefix = prefix var eventsToStore []func() + var affectedData *resourceAffectedPeersData err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { network, err := transaction.GetNetworkByID(ctx, store.LockingStrengthUpdate, resource.AccountID, resource.NetworkID) if err != nil { @@ -232,6 +241,15 @@ func (m *managerImpl) UpdateResource(ctx context.Context, userID string, resourc return fmt.Errorf("failed to get network resource: %w", err) } + oldGroups, err := m.groupsManager.GetResourceGroupsInTransaction(ctx, transaction, store.LockingStrengthNone, oldResource.AccountID, oldResource.ID) + if err != nil { + return fmt.Errorf("failed to get old resource groups: %w", err) + } + var oldGroupIDs []string + for _, g := range oldGroups { + oldGroupIDs = append(oldGroupIDs, g.ID) + } + err = transaction.SaveNetworkResource(ctx, resource) if err != nil { return fmt.Errorf("failed to save network resource: %w", err) @@ -247,6 +265,11 @@ func (m *managerImpl) UpdateResource(ctx context.Context, userID string, resourc m.accountManager.StoreEvent(ctx, userID, resource.ID, resource.AccountID, activity.NetworkResourceUpdated, resource.EventMeta(network)) }) + affectedData, err = loadResourceAffectedPeersData(ctx, transaction, resource.AccountID, resource.NetworkID, append(resource.GroupIDs, oldGroupIDs...)) + if err != nil { + log.WithContext(ctx).Errorf("failed to load affected peers data: %v", err) + } + err = transaction.IncrementNetworkSerial(ctx, resource.AccountID) if err != nil { return fmt.Errorf("failed to increment network serial: %w", err) @@ -270,7 +293,9 @@ func (m *managerImpl) UpdateResource(ctx context.Context, userID string, resourc } }() - go m.accountManager.UpdateAccountPeers(ctx, resource.AccountID) + if affectedPeerIDs := m.resolveResourceAffectedPeers(ctx, resource.AccountID, affectedData); len(affectedPeerIDs) > 0 { + go m.accountManager.UpdateAffectedPeers(ctx, resource.AccountID, affectedPeerIDs) + } return resource, nil } @@ -331,7 +356,22 @@ func (m *managerImpl) DeleteResource(ctx context.Context, accountID, userID, net } var events []func() + var affectedData *resourceAffectedPeersData err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { + groups, err := m.groupsManager.GetResourceGroupsInTransaction(ctx, transaction, store.LockingStrengthNone, accountID, resourceID) + if err != nil { + return fmt.Errorf("failed to get resource groups: %w", err) + } + var resourceGroupIDs []string + for _, g := range groups { + resourceGroupIDs = append(resourceGroupIDs, g.ID) + } + + affectedData, err = loadResourceAffectedPeersData(ctx, transaction, accountID, networkID, resourceGroupIDs) + if err != nil { + log.WithContext(ctx).Errorf("failed to load affected peers data: %v", err) + } + events, err = m.DeleteResourceInTransaction(ctx, transaction, accountID, userID, networkID, resourceID) if err != nil { return fmt.Errorf("failed to delete resource: %w", err) @@ -352,7 +392,9 @@ func (m *managerImpl) DeleteResource(ctx context.Context, accountID, userID, net event() } - go m.accountManager.UpdateAccountPeers(ctx, accountID) + if affectedPeerIDs := m.resolveResourceAffectedPeers(ctx, accountID, affectedData); len(affectedPeerIDs) > 0 { + go m.accountManager.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } return nil } @@ -399,6 +441,127 @@ func (m *managerImpl) DeleteResourceInTransaction(ctx context.Context, transacti return eventsToStore, nil } +// resourceAffectedPeersData holds data loaded inside a transaction for affected peer resolution. +type resourceAffectedPeersData struct { + resourceGroupIDs []string + policies []*nbtypes.Policy + routerPeerGroups []string + routerDirectPeers []string +} + +// loadResourceAffectedPeersData loads the data needed to determine affected peers within a transaction. +func loadResourceAffectedPeersData(ctx context.Context, transaction store.Store, accountID, networkID string, resourceGroupIDs []string) (*resourceAffectedPeersData, error) { + if len(resourceGroupIDs) == 0 { + return nil, nil + } + + policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + if err != nil { + return nil, fmt.Errorf("failed to get policies: %w", err) + } + + routers, err := transaction.GetNetworkRoutersByNetID(ctx, store.LockingStrengthNone, accountID, networkID) + if err != nil { + return nil, fmt.Errorf("failed to get routers: %w", err) + } + + var routerPeerGroups []string + var routerDirectPeers []string + for _, router := range routers { + if !router.Enabled { + continue + } + routerPeerGroups = append(routerPeerGroups, router.PeerGroups...) + if router.Peer != "" { + routerDirectPeers = append(routerDirectPeers, router.Peer) + } + } + + return &resourceAffectedPeersData{ + resourceGroupIDs: resourceGroupIDs, + policies: policies, + routerPeerGroups: routerPeerGroups, + routerDirectPeers: routerDirectPeers, + }, nil +} + +// resolveResourceAffectedPeers computes affected peer IDs from preloaded data outside the transaction. +func (m *managerImpl) resolveResourceAffectedPeers(ctx context.Context, accountID string, data *resourceAffectedPeersData) []string { + if data == nil { + return nil + } + + groupSet := make(map[string]struct{}) + var directPeerIDs []string + + destSet := make(map[string]struct{}, len(data.resourceGroupIDs)) + for _, gID := range data.resourceGroupIDs { + destSet[gID] = struct{}{} + } + + for _, policy := range data.policies { + if policy == nil || !policy.Enabled { + continue + } + for _, rule := range policy.Rules { + if rule == nil || !rule.Enabled { + continue + } + referencesResource := false + for _, gID := range rule.Destinations { + if _, ok := destSet[gID]; ok { + referencesResource = true + break + } + } + if !referencesResource { + continue + } + for _, gID := range rule.Sources { + groupSet[gID] = struct{}{} + } + if rule.SourceResource.Type == nbtypes.ResourceTypePeer && rule.SourceResource.ID != "" { + directPeerIDs = append(directPeerIDs, rule.SourceResource.ID) + } + } + } + + for _, gID := range data.routerPeerGroups { + groupSet[gID] = struct{}{} + } + directPeerIDs = append(directPeerIDs, data.routerDirectPeers...) + + if len(groupSet) == 0 && len(directPeerIDs) == 0 { + return nil + } + + groupIDs := make([]string, 0, len(groupSet)) + for gID := range groupSet { + groupIDs = append(groupIDs, gID) + } + + peerIDs, err := m.store.GetPeerIDsByGroups(ctx, accountID, groupIDs) + if err != nil { + log.WithContext(ctx).Errorf("failed to resolve peer IDs: %v", err) + return nil + } + + if len(directPeerIDs) > 0 { + seen := make(map[string]struct{}, len(peerIDs)) + for _, id := range peerIDs { + seen[id] = struct{}{} + } + for _, id := range directPeerIDs { + if _, exists := seen[id]; !exists { + peerIDs = append(peerIDs, id) + seen[id] = struct{}{} + } + } + } + + return peerIDs, nil +} + func NewManagerMock() Manager { return &mockManager{} } diff --git a/management/server/networks/routers/manager.go b/management/server/networks/routers/manager.go index 82cac424a..7ee9dc281 100644 --- a/management/server/networks/routers/manager.go +++ b/management/server/networks/routers/manager.go @@ -6,6 +6,7 @@ import ( "fmt" "github.com/rs/xid" + log "github.com/sirupsen/logrus" "github.com/netbirdio/netbird/management/server/account" "github.com/netbirdio/netbird/management/server/activity" @@ -15,6 +16,7 @@ import ( "github.com/netbirdio/netbird/management/server/permissions/modules" "github.com/netbirdio/netbird/management/server/permissions/operations" "github.com/netbirdio/netbird/management/server/store" + nbtypes "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/shared/management/status" ) @@ -89,6 +91,7 @@ func (m *managerImpl) CreateRouter(ctx context.Context, userID string, router *t } var network *networkTypes.Network + var affectedData *routerAffectedPeersData err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { network, err = transaction.GetNetworkByID(ctx, store.LockingStrengthNone, router.AccountID, router.NetworkID) if err != nil { @@ -111,6 +114,11 @@ func (m *managerImpl) CreateRouter(ctx context.Context, userID string, router *t return fmt.Errorf("failed to increment network serial: %w", err) } + affectedData, err = loadRouterAffectedPeersData(ctx, transaction, router.AccountID, router.NetworkID, router.PeerGroups, router.Peer) + if err != nil { + log.WithContext(ctx).Errorf("failed to load affected peers data: %v", err) + } + return nil }) if err != nil { @@ -119,7 +127,9 @@ func (m *managerImpl) CreateRouter(ctx context.Context, userID string, router *t m.accountManager.StoreEvent(ctx, userID, router.ID, router.AccountID, activity.NetworkRouterCreated, router.EventMeta(network)) - go m.accountManager.UpdateAccountPeers(ctx, router.AccountID) + if affectedPeerIDs := m.resolveRouterAffectedPeers(ctx, router.AccountID, affectedData); len(affectedPeerIDs) > 0 { + go m.accountManager.UpdateAffectedPeers(ctx, router.AccountID, affectedPeerIDs) + } return router, nil } @@ -155,6 +165,7 @@ func (m *managerImpl) UpdateRouter(ctx context.Context, userID string, router *t } var network *networkTypes.Network + var affectedData *routerAffectedPeersData err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { network, err = transaction.GetNetworkByID(ctx, store.LockingStrengthNone, router.AccountID, router.NetworkID) if err != nil { @@ -165,6 +176,16 @@ func (m *managerImpl) UpdateRouter(ctx context.Context, userID string, router *t return status.NewRouterNotPartOfNetworkError(router.ID, router.NetworkID) } + allPeerGroups := router.PeerGroups + directPeers := []string{router.Peer} + oldRouter, err := transaction.GetNetworkRouterByID(ctx, store.LockingStrengthNone, router.AccountID, router.ID) + if err == nil { + allPeerGroups = append(allPeerGroups, oldRouter.PeerGroups...) + if oldRouter.Peer != "" { + directPeers = append(directPeers, oldRouter.Peer) + } + } + err = transaction.SaveNetworkRouter(ctx, router) if err != nil { return fmt.Errorf("failed to update network router: %w", err) @@ -175,6 +196,11 @@ func (m *managerImpl) UpdateRouter(ctx context.Context, userID string, router *t return fmt.Errorf("failed to increment network serial: %w", err) } + affectedData, err = loadRouterAffectedPeersData(ctx, transaction, router.AccountID, router.NetworkID, allPeerGroups, directPeers...) + if err != nil { + log.WithContext(ctx).Errorf("failed to load affected peers data: %v", err) + } + return nil }) if err != nil { @@ -183,7 +209,9 @@ func (m *managerImpl) UpdateRouter(ctx context.Context, userID string, router *t m.accountManager.StoreEvent(ctx, userID, router.ID, router.AccountID, activity.NetworkRouterUpdated, router.EventMeta(network)) - go m.accountManager.UpdateAccountPeers(ctx, router.AccountID) + if affectedPeerIDs := m.resolveRouterAffectedPeers(ctx, router.AccountID, affectedData); len(affectedPeerIDs) > 0 { + go m.accountManager.UpdateAffectedPeers(ctx, router.AccountID, affectedPeerIDs) + } return router, nil } @@ -198,7 +226,19 @@ func (m *managerImpl) DeleteRouter(ctx context.Context, accountID, userID, netwo } var event func() + var affectedData *routerAffectedPeersData err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { + router, err := transaction.GetNetworkRouterByID(ctx, store.LockingStrengthNone, accountID, routerID) + if err != nil { + return fmt.Errorf("failed to get router: %w", err) + } + + // load before delete so group memberships are still present + affectedData, err = loadRouterAffectedPeersData(ctx, transaction, accountID, networkID, router.PeerGroups, router.Peer) + if err != nil { + log.WithContext(ctx).Errorf("failed to load affected peers data: %v", err) + } + event, err = m.DeleteRouterInTransaction(ctx, transaction, accountID, userID, networkID, routerID) if err != nil { return fmt.Errorf("failed to delete network router: %w", err) @@ -217,7 +257,9 @@ func (m *managerImpl) DeleteRouter(ctx context.Context, accountID, userID, netwo event() - go m.accountManager.UpdateAccountPeers(ctx, accountID) + if affectedPeerIDs := m.resolveRouterAffectedPeers(ctx, accountID, affectedData); len(affectedPeerIDs) > 0 { + go m.accountManager.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } return nil } @@ -249,6 +291,136 @@ func (m *managerImpl) DeleteRouterInTransaction(ctx context.Context, transaction return event, nil } +// routerAffectedPeersData holds data loaded inside a transaction for affected peer resolution. +type routerAffectedPeersData struct { + routerPeerGroups []string + directPeerIDs []string + resourceGroupIDs []string + policies []*nbtypes.Policy +} + +// loadRouterAffectedPeersData loads the data needed to determine affected peers within a transaction. +func loadRouterAffectedPeersData(ctx context.Context, transaction store.Store, accountID, networkID string, routerPeerGroups []string, directPeers ...string) (*routerAffectedPeersData, error) { + var directPeerIDs []string + for _, p := range directPeers { + if p != "" { + directPeerIDs = append(directPeerIDs, p) + } + } + + if len(routerPeerGroups) == 0 && len(directPeerIDs) == 0 { + return nil, nil + } + + resources, err := transaction.GetNetworkResourcesByNetID(ctx, store.LockingStrengthNone, accountID, networkID) + if err != nil { + return nil, fmt.Errorf("failed to get network resources: %w", err) + } + + var resourceGroupIDs []string + for _, resource := range resources { + if !resource.Enabled { + continue + } + groups, err := transaction.GetResourceGroups(ctx, store.LockingStrengthNone, accountID, resource.ID) + if err != nil { + return nil, fmt.Errorf("failed to get groups for resource %s: %w", resource.ID, err) + } + for _, g := range groups { + resourceGroupIDs = append(resourceGroupIDs, g.ID) + } + } + + var policies []*nbtypes.Policy + if len(resourceGroupIDs) > 0 { + policies, err = transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + if err != nil { + return nil, fmt.Errorf("failed to get policies: %w", err) + } + } + + return &routerAffectedPeersData{ + routerPeerGroups: routerPeerGroups, + directPeerIDs: directPeerIDs, + resourceGroupIDs: resourceGroupIDs, + policies: policies, + }, nil +} + +// resolveRouterAffectedPeers computes affected peer IDs from preloaded data outside the transaction. +func (m *managerImpl) resolveRouterAffectedPeers(ctx context.Context, accountID string, data *routerAffectedPeersData) []string { + if data == nil { + return nil + } + + groupSet := make(map[string]struct{}) + + for _, gID := range data.routerPeerGroups { + groupSet[gID] = struct{}{} + } + + if len(data.resourceGroupIDs) > 0 { + destSet := make(map[string]struct{}, len(data.resourceGroupIDs)) + for _, gID := range data.resourceGroupIDs { + destSet[gID] = struct{}{} + } + + for _, policy := range data.policies { + if policy == nil || !policy.Enabled { + continue + } + for _, rule := range policy.Rules { + if rule == nil || !rule.Enabled { + continue + } + referencesResource := false + for _, gID := range rule.Destinations { + if _, ok := destSet[gID]; ok { + referencesResource = true + break + } + } + if !referencesResource { + continue + } + for _, gID := range rule.Sources { + groupSet[gID] = struct{}{} + } + } + } + } + + if len(groupSet) == 0 && len(data.directPeerIDs) == 0 { + return nil + } + + groupIDs := make([]string, 0, len(groupSet)) + for gID := range groupSet { + groupIDs = append(groupIDs, gID) + } + + peerIDs, err := m.store.GetPeerIDsByGroups(ctx, accountID, groupIDs) + if err != nil { + log.WithContext(ctx).Errorf("failed to resolve peer IDs: %v", err) + return nil + } + + if len(data.directPeerIDs) > 0 { + seen := make(map[string]struct{}, len(peerIDs)) + for _, id := range peerIDs { + seen[id] = struct{}{} + } + for _, id := range data.directPeerIDs { + if _, exists := seen[id]; !exists { + peerIDs = append(peerIDs, id) + seen[id] = struct{}{} + } + } + } + + return peerIDs +} + func NewManagerMock() Manager { return &mockManager{} } From 6b4d4076f4b834590684dc8a88850861b47ef2eb Mon Sep 17 00:00:00 2001 From: pascal Date: Mon, 4 May 2026 15:16:59 +0200 Subject: [PATCH 07/28] extend tests --- management/server/affected_peers_test.go | 1214 +++++++++++++++++++--- 1 file changed, 1072 insertions(+), 142 deletions(-) diff --git a/management/server/affected_peers_test.go b/management/server/affected_peers_test.go index bf48b0e4a..0004fe5c1 100644 --- a/management/server/affected_peers_test.go +++ b/management/server/affected_peers_test.go @@ -15,6 +15,7 @@ import ( routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" networkTypes "github.com/netbirdio/netbird/management/server/networks/types" nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/management/server/posture" "github.com/netbirdio/netbird/management/server/store" "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/route" @@ -144,6 +145,30 @@ func TestCollectGroupChange_PolicyWithDirectPeerResource(t *testing.T) { assert.Contains(t, directPeers, peerIDs[3]) } +func TestCollectGroupChange_PolicyWithNonPeerResource_NoDirectPeers(t *testing.T) { + manager, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[0]}, + SourceResource: types.Resource{ID: "some-domain", Type: types.ResourceTypeDomain}, + Destinations: []string{groupIDs[1]}, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) + assert.Contains(t, groups, groupIDs[0]) + assert.Contains(t, groups, groupIDs[1]) + assert.Empty(t, directPeers, "non-peer resources should not produce direct peer IDs") +} + func TestCollectGroupChange_RouteLinked(t *testing.T) { manager, s, accountID, _, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() @@ -277,6 +302,35 @@ func TestCollectGroupChange_NetworkRouterLinked(t *testing.T) { assert.Empty(t, directPeers) } +func TestCollectGroupChange_NetworkRouterPeerOnlyNoGroups(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + net1 := &networkTypes.Network{ + ID: "net-peer-only", + AccountID: accountID, + Name: "peer-only-network", + } + err := manager.Store.SaveNetwork(ctx, net1) + require.NoError(t, err) + + // Router with only a direct peer, no PeerGroups + err = manager.Store.SaveNetworkRouter(ctx, &routerTypes.NetworkRouter{ + ID: "router-peer-only", + NetworkID: net1.ID, + AccountID: accountID, + Peer: peerIDs[4], + }) + require.NoError(t, err) + + // None of the groups should match since router has no PeerGroups + for i := 0; i < 5; i++ { + groups, directPeers := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[i]}) + assert.Empty(t, groups, "group%d should not match router with only direct peer", i) + assert.Empty(t, directPeers, "group%d should not produce direct peers", i) + } +} + func TestCollectGroupChange_MultipleEntities(t *testing.T) { manager, s, accountID, _, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() @@ -329,6 +383,51 @@ func TestCollectGroupChange_MultipleEntities(t *testing.T) { assert.Empty(t, directPeers) } +func TestCollectGroupChange_MultipleNameServerGroups_OnlyLinkedAffected(t *testing.T) { + manager, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Create two nameserver groups using different groups + _, err := manager.CreateNameServerGroup(ctx, accountID, "ns-a", "NS-A", + []nbdns.NameServer{{ + IP: netip.MustParseAddr("1.1.1.1"), + NSType: nbdns.UDPNameServerType, + Port: nbdns.DefaultDNSPort, + }}, + []string{groupIDs[0]}, + true, nil, true, userID, false, + ) + require.NoError(t, err) + + _, err = manager.CreateNameServerGroup(ctx, accountID, "ns-b", "NS-B", + []nbdns.NameServer{{ + IP: netip.MustParseAddr("8.8.8.8"), + NSType: nbdns.UDPNameServerType, + Port: nbdns.DefaultDNSPort, + }}, + []string{groupIDs[2]}, + true, nil, true, userID, false, + ) + require.NoError(t, err) + + // Changing group0 should only find group0 (from ns-a), not group2 (from ns-b) + groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) + assert.Contains(t, groups, groupIDs[0]) + assert.NotContains(t, groups, groupIDs[2]) + + groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[2]}) + assert.Contains(t, groups, groupIDs[2]) + assert.NotContains(t, groups, groupIDs[0]) + + // Unrelated group + groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[4]}) + assert.Empty(t, groups) +} + +// --------------------------------------------------------------------------- +// collectPolicyAffectedGroupsAndPeers unit tests +// --------------------------------------------------------------------------- + func TestCollectPolicyAffectedGroups_Basic(t *testing.T) { policy := &types.Policy{ Rules: []*types.PolicyRule{ @@ -376,6 +475,47 @@ func TestCollectPolicyAffectedGroups_MultipleRules(t *testing.T) { assert.ElementsMatch(t, []string{"g1", "g2", "g3", "g4"}, groups) } +func TestCollectPolicyAffectedGroups_MultiplePolicies(t *testing.T) { + old := &types.Policy{ + Rules: []*types.PolicyRule{ + {Sources: []string{"g1"}, Destinations: []string{"g2"}}, + }, + } + new := &types.Policy{ + Rules: []*types.PolicyRule{ + {Sources: []string{"g3"}, Destinations: []string{"g4"}}, + }, + } + groups, _ := collectPolicyAffectedGroupsAndPeers(new, old) + assert.ElementsMatch(t, []string{"g1", "g2", "g3", "g4"}, groups) +} + +func TestCollectPolicyAffectedGroups_EmptyRules(t *testing.T) { + policy := &types.Policy{Rules: []*types.PolicyRule{}} + groups, directPeers := collectPolicyAffectedGroupsAndPeers(policy) + assert.Empty(t, groups) + assert.Empty(t, directPeers) +} + +func TestCollectPolicyAffectedGroups_NonPeerResource(t *testing.T) { + policy := &types.Policy{ + Rules: []*types.PolicyRule{ + { + Sources: []string{"g1"}, + SourceResource: types.Resource{ID: "domain-1", Type: types.ResourceTypeDomain}, + Destinations: []string{"g2"}, + }, + }, + } + groups, directPeers := collectPolicyAffectedGroupsAndPeers(policy) + assert.ElementsMatch(t, []string{"g1", "g2"}, groups) + assert.Empty(t, directPeers, "domain resource type should not produce direct peer IDs") +} + +// --------------------------------------------------------------------------- +// collectRouteAffectedGroupsAndPeers unit tests +// --------------------------------------------------------------------------- + func TestCollectRouteAffectedGroups_Basic(t *testing.T) { r := &route.Route{ Groups: []string{"g1"}, @@ -403,6 +543,105 @@ func TestCollectRouteAffectedGroups_NilRoute(t *testing.T) { assert.Nil(t, directPeers) } +func TestCollectRouteAffectedGroups_MultipleRoutes(t *testing.T) { + old := &route.Route{ + Groups: []string{"g1"}, + Peer: "p1", + } + new := &route.Route{ + Groups: []string{"g2"}, + PeerGroups: []string{"g3"}, + } + groups, directPeers := collectRouteAffectedGroupsAndPeers(new, old) + assert.ElementsMatch(t, []string{"g1", "g2", "g3"}, groups) + assert.ElementsMatch(t, []string{"p1"}, directPeers) +} + +// --------------------------------------------------------------------------- +// policyReferencesGroups / routeReferencesGroups / routerReferencesGroups +// --------------------------------------------------------------------------- + +func TestPolicyReferencesGroups(t *testing.T) { + policy := &types.Policy{ + Rules: []*types.PolicyRule{ + { + Sources: []string{"g1", "g2"}, + Destinations: []string{"g3"}, + }, + }, + } + + tests := []struct { + name string + groupSet map[string]struct{} + want bool + }{ + {"matches source", map[string]struct{}{"g1": {}}, true}, + {"matches destination", map[string]struct{}{"g3": {}}, true}, + {"no match", map[string]struct{}{"g4": {}}, false}, + {"empty set", map[string]struct{}{}, false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := policyReferencesGroups(policy, tt.groupSet) + assert.Equal(t, tt.want, got) + }) + } +} + +func TestRouteReferencesGroups(t *testing.T) { + r := &route.Route{ + Groups: []string{"g1"}, + PeerGroups: []string{"g2"}, + AccessControlGroups: []string{"g3"}, + } + + tests := []struct { + name string + groupSet map[string]struct{} + want bool + }{ + {"matches groups", map[string]struct{}{"g1": {}}, true}, + {"matches peerGroups", map[string]struct{}{"g2": {}}, true}, + {"matches accessControl", map[string]struct{}{"g3": {}}, true}, + {"no match", map[string]struct{}{"g4": {}}, false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := routeReferencesGroups(r, tt.groupSet) + assert.Equal(t, tt.want, got) + }) + } +} + +func TestRouterReferencesGroups(t *testing.T) { + router := &routerTypes.NetworkRouter{ + PeerGroups: []string{"g1", "g2"}, + } + + tests := []struct { + name string + groupSet map[string]struct{} + want bool + }{ + {"matches", map[string]struct{}{"g1": {}}, true}, + {"no match", map[string]struct{}{"g3": {}}, false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := routerReferencesGroups(router, tt.groupSet) + assert.Equal(t, tt.want, got) + }) + } +} + +// --------------------------------------------------------------------------- +// resolvePeerIDs tests +// --------------------------------------------------------------------------- + func TestResolvePeerIDs_GroupsOnly(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() @@ -436,6 +675,10 @@ func TestResolvePeerIDs_EmptyInputs(t *testing.T) { assert.Empty(t, result) } +// --------------------------------------------------------------------------- +// resolveAffectedPeersForPeerChanges tests +// --------------------------------------------------------------------------- + func TestResolveAffectedPeers_NoPoliciesOrRoutes(t *testing.T) { manager, s, accountID, peerIDs, _ := setupAffectedPeersTest(t) ctx := context.Background() @@ -553,6 +796,38 @@ func TestResolveAffectedPeers_RouteWithDirectPeer(t *testing.T) { assert.ElementsMatch(t, []string{peerIDs[1], peerIDs[4]}, result) } +func TestResolveAffectedPeers_RouteWithAccessControlGroups(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + _, err := manager.CreateRoute(ctx, accountID, + netip.MustParsePrefix("10.7.0.0/24"), + route.IPv4Network, + nil, + "", + []string{groupIDs[0]}, + "acl route", + "aclnet", + false, + 9999, + []string{groupIDs[1]}, + []string{groupIDs[2]}, + true, + userID, + false, + false, + ) + require.NoError(t, err) + + // peer2 is only in AccessControlGroups, still should be affected + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[2]}) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1], peerIDs[2]}, result) + + // peer3 is unrelated + result = manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[3]}) + assert.Empty(t, result) +} + func TestResolveAffectedPeers_NetworkRouter(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() @@ -737,90 +1012,42 @@ func TestResolveAffectedPeers_EmptyChangedPeers(t *testing.T) { assert.Empty(t, result) } -func TestAffectedPeers_GroupUpdateOnlyAffectsLinkedPeers(t *testing.T) { - manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) +func TestResolveAffectedPeers_NoDuplicates(t *testing.T) { + manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() - accountID := account.Id - policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + err := manager.GroupAddPeer(ctx, accountID, groupIDs[1], peerIDs[0]) + require.NoError(t, err) + err = manager.GroupAddPeer(ctx, accountID, groupIDs[2], peerIDs[0]) require.NoError(t, err) - for _, p := range policies { - err := manager.Store.DeletePolicy(ctx, accountID, p.ID) - require.NoError(t, err) - } - - for _, g := range []*types.Group{ - {ID: "ap-grpA", Name: "AP-A", Peers: []string{peer1.ID}}, - {ID: "ap-grpB", Name: "AP-B", Peers: []string{peer2.ID}}, - {ID: "ap-grpC", Name: "AP-C", Peers: []string{peer3.ID}}, - } { - err := manager.CreateGroup(ctx, accountID, userID, g) - require.NoError(t, err) - } _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ { - Enabled: true, - Sources: []string{"ap-grpA"}, - Destinations: []string{"ap-grpB"}, - Bidirectional: true, - Action: types.PolicyTrafficActionAccept, + Enabled: true, + Sources: []string{groupIDs[0], groupIDs[1]}, + Destinations: []string{groupIDs[2]}, + Action: types.PolicyTrafficActionAccept, }, }, }, true) require.NoError(t, err) - updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) - updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) - updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) - t.Cleanup(func() { - updateManager.CloseChannel(ctx, peer1.ID) - updateManager.CloseChannel(ctx, peer2.ID) - updateManager.CloseChannel(ctx, peer3.ID) - }) - - result := manager.resolveAffectedPeersForPeerChanges(ctx, manager.Store, accountID, []string{peer1.ID}) - assert.ElementsMatch(t, []string{peer1.ID, peer2.ID}, result) - - // Adding peer3 to grpA makes it part of the policy, so all 3 peers get updated - t.Run("group change updates all peers in policy groups", func(t *testing.T) { - done := make(chan struct{}) - go func() { - peerShouldReceiveUpdate(t, updMsg1) - peerShouldReceiveUpdate(t, updMsg2) - peerShouldReceiveUpdate(t, updMsg3) - close(done) - }() - - err := manager.UpdateGroup(ctx, accountID, userID, &types.Group{ - ID: "ap-grpA", - Name: "AP-A", - Peers: []string{peer1.ID, peer3.ID}, - }) - assert.NoError(t, err) - - select { - case <-done: - case <-time.After(peerUpdateTimeout): - t.Error("timeout") - } - }) - - _ = updMsg1 - _ = updMsg2 - _ = updMsg3 -} - -func TestAffectedPeers_UnlinkedGroupChange_NoUpdates(t *testing.T) { - manager, s, accountID, peerIDs, _ := setupAffectedPeersTest(t) - ctx := context.Background() - result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) - assert.Empty(t, result) + count := 0 + for _, id := range result { + if id == peerIDs[0] { + count++ + } + } + assert.Equal(t, 1, count, "peer0 should appear exactly once") } +// --------------------------------------------------------------------------- +// Posture check affected peers tests +// --------------------------------------------------------------------------- + func TestCollectPostureCheckAffected_NoMatch(t *testing.T) { _, s, accountID, _, _ := setupAffectedPeersTest(t) ctx := context.Background() @@ -830,6 +1057,48 @@ func TestCollectPostureCheckAffected_NoMatch(t *testing.T) { assert.Empty(t, directPeers) } +func TestCollectPostureCheckAffected_LinkedToPolicy(t *testing.T) { + manager, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Create the posture check in the store so the policy validation keeps the reference. + err := s.SavePostureChecks(ctx, &posture.Checks{ + ID: "pc-1", + Name: "test-posture-check", + AccountID: accountID, + }) + require.NoError(t, err) + + policy, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + SourcePostureChecks: []string{"pc-1"}, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{groupIDs[0]}, + Destinations: []string{groupIDs[1]}, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + _ = policy + + groups, directPeers := collectPostureCheckAffectedGroupsAndPeers(ctx, s, accountID, "pc-1") + assert.Contains(t, groups, groupIDs[0]) + assert.Contains(t, groups, groupIDs[1]) + assert.Empty(t, directPeers) + + // Different posture check ID should not match + groups, directPeers = collectPostureCheckAffectedGroupsAndPeers(ctx, s, accountID, "pc-other") + assert.Empty(t, groups) + assert.Empty(t, directPeers) +} + +// --------------------------------------------------------------------------- +// Isolation tests: verify peers NOT in any relevant entity are NOT affected +// --------------------------------------------------------------------------- + func TestAffectedPeers_IsolatedPolicies(t *testing.T) { manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) ctx := context.Background() @@ -924,116 +1193,777 @@ func TestAffectedPeers_IsolatedRouteAndPolicy(t *testing.T) { assert.NotContains(t, result, peerIDs[1]) } -func TestResolveAffectedPeers_NoDuplicates(t *testing.T) { - manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) - ctx := context.Background() +// --------------------------------------------------------------------------- +// Integration tests with update channels (peerShouldReceiveUpdate / peerShouldNotReceiveUpdate) +// --------------------------------------------------------------------------- - err := manager.GroupAddPeer(ctx, accountID, groupIDs[1], peerIDs[0]) - require.NoError(t, err) - err = manager.GroupAddPeer(ctx, accountID, groupIDs[2], peerIDs[0]) +func TestAffectedPeers_GroupUpdateOnlyAffectsLinkedPeers(t *testing.T) { + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) require.NoError(t, err) + for _, p := range policies { + err := manager.Store.DeletePolicy(ctx, accountID, p.ID) + require.NoError(t, err) + } + + for _, g := range []*types.Group{ + {ID: "ap-grpA", Name: "AP-A", Peers: []string{peer1.ID}}, + {ID: "ap-grpB", Name: "AP-B", Peers: []string{peer2.ID}}, + {ID: "ap-grpC", Name: "AP-C", Peers: []string{peer3.ID}}, + } { + err := manager.CreateGroup(ctx, accountID, userID, g) + require.NoError(t, err) + } _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ Enabled: true, Rules: []*types.PolicyRule{ { - Enabled: true, - Sources: []string{groupIDs[0], groupIDs[1]}, - Destinations: []string{groupIDs[2]}, - Action: types.PolicyTrafficActionAccept, + Enabled: true, + Sources: []string{"ap-grpA"}, + Destinations: []string{"ap-grpB"}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, }, }, }, true) require.NoError(t, err) - result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) - count := 0 - for _, id := range result { - if id == peerIDs[0] { - count++ + updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) + updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) + updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + result := manager.resolveAffectedPeersForPeerChanges(ctx, manager.Store, accountID, []string{peer1.ID}) + assert.ElementsMatch(t, []string{peer1.ID, peer2.ID}, result) + + // Adding peer3 to grpA makes it part of the policy, so all 3 peers get updated + t.Run("group change updates all peers in policy groups", func(t *testing.T) { + done := make(chan struct{}) + go func() { + peerShouldReceiveUpdate(t, updMsg1) + peerShouldReceiveUpdate(t, updMsg2) + peerShouldReceiveUpdate(t, updMsg3) + close(done) + }() + + err := manager.UpdateGroup(ctx, accountID, userID, &types.Group{ + ID: "ap-grpA", + Name: "AP-A", + Peers: []string{peer1.ID, peer3.ID}, + }) + assert.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout") } - } - assert.Equal(t, 1, count, "peer0 should appear exactly once") + }) + + _ = updMsg1 + _ = updMsg2 + _ = updMsg3 } -func TestPolicyReferencesGroups(t *testing.T) { - policy := &types.Policy{ +func TestAffectedPeers_UnlinkedGroupChange_NoUpdates(t *testing.T) { + manager, s, accountID, peerIDs, _ := setupAffectedPeersTest(t) + ctx := context.Background() + + result := manager.resolveAffectedPeersForPeerChanges(ctx, s, accountID, []string{peerIDs[0]}) + assert.Empty(t, result) +} + +// TestAffectedPeers_PolicyChange_UnrelatedPeerNoUpdate verifies that creating/deleting a +// policy only sends updates to peers in the policy's groups, not to unrelated peers. +func TestAffectedPeers_PolicyChange_UnrelatedPeerNoUpdate(t *testing.T) { + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + err := manager.Store.DeletePolicy(ctx, accountID, p.ID) + require.NoError(t, err) + } + + for _, g := range []*types.Group{ + {ID: "pol-grpA", Name: "Pol-A", Peers: []string{peer1.ID}}, + {ID: "pol-grpB", Name: "Pol-B", Peers: []string{peer2.ID}}, + {ID: "pol-grpC", Name: "Pol-C", Peers: []string{peer3.ID}}, + } { + err := manager.CreateGroup(ctx, accountID, userID, g) + require.NoError(t, err) + } + + updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) + updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) + updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + // Create policy linking only peer1 (grpA) <-> peer2 (grpB). Peer3 should not receive update. + t.Run("create policy only affects linked peers", func(t *testing.T) { + done := make(chan struct{}) + go func() { + peerShouldReceiveUpdate(t, updMsg1) + peerShouldReceiveUpdate(t, updMsg2) + peerShouldNotReceiveUpdate(t, updMsg3) + close(done) + }() + + _, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{"pol-grpA"}, + Destinations: []string{"pol-grpB"}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + assert.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout") + } + }) +} + +// TestAffectedPeers_RouteChange_UnrelatedPeerNoUpdate verifies that creating a route +// only sends updates to peers in the route's groups, not to unrelated peers. +func TestAffectedPeers_RouteChange_UnrelatedPeerNoUpdate(t *testing.T) { + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + err := manager.Store.DeletePolicy(ctx, accountID, p.ID) + require.NoError(t, err) + } + + for _, g := range []*types.Group{ + {ID: "rt-grpA", Name: "Rt-A", Peers: []string{peer1.ID}}, + {ID: "rt-grpB", Name: "Rt-B", Peers: []string{peer2.ID}}, + {ID: "rt-grpC", Name: "Rt-C", Peers: []string{peer3.ID}}, + } { + err := manager.CreateGroup(ctx, accountID, userID, g) + require.NoError(t, err) + } + + updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) + updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) + updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + // Create route with peer groups grpA and distribution group grpB. Peer3 should not get update. + t.Run("create route only affects linked peers", func(t *testing.T) { + done := make(chan struct{}) + go func() { + peerShouldReceiveUpdate(t, updMsg1) + peerShouldReceiveUpdate(t, updMsg2) + peerShouldNotReceiveUpdate(t, updMsg3) + close(done) + }() + + _, err := manager.CreateRoute(ctx, accountID, + netip.MustParsePrefix("10.10.0.0/24"), + route.IPv4Network, + nil, + "", + []string{"rt-grpA"}, + "test route", + "routenoaffect", + false, + 9999, + []string{"rt-grpB"}, + nil, + true, + userID, + false, + false, + ) + assert.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout") + } + }) +} + +// TestAffectedPeers_NameServerChange_UnrelatedPeerNoUpdate verifies that creating a +// nameserver group only sends updates to peers in its groups, not to unrelated peers. +func TestAffectedPeers_NameServerChange_UnrelatedPeerNoUpdate(t *testing.T) { + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + err := manager.Store.DeletePolicy(ctx, accountID, p.ID) + require.NoError(t, err) + } + + for _, g := range []*types.Group{ + {ID: "ns-grpA", Name: "NS-A", Peers: []string{peer1.ID}}, + {ID: "ns-grpB", Name: "NS-B", Peers: []string{peer2.ID}}, + } { + err := manager.CreateGroup(ctx, accountID, userID, g) + require.NoError(t, err) + } + + updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) + updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) + updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + // Create NS group using only grpA. peer2 and peer3 should not get update. + t.Run("create nameserver group only affects linked peers", func(t *testing.T) { + done := make(chan struct{}) + go func() { + peerShouldReceiveUpdate(t, updMsg1) + peerShouldNotReceiveUpdate(t, updMsg2) + peerShouldNotReceiveUpdate(t, updMsg3) + close(done) + }() + + _, err := manager.CreateNameServerGroup(ctx, accountID, "ns-unrelated", "NS Unrelated", + []nbdns.NameServer{{ + IP: netip.MustParseAddr("1.1.1.1"), + NSType: nbdns.UDPNameServerType, + Port: nbdns.DefaultDNSPort, + }}, + []string{"ns-grpA"}, + true, nil, true, userID, false, + ) + assert.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout") + } + }) +} + +// TestAffectedPeers_DNSSettingsChange_UnrelatedPeerNoUpdate verifies that changing DNS +// settings only sends updates to peers in the affected groups, not to unrelated peers. +func TestAffectedPeers_DNSSettingsChange_UnrelatedPeerNoUpdate(t *testing.T) { + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + err := manager.Store.DeletePolicy(ctx, accountID, p.ID) + require.NoError(t, err) + } + + for _, g := range []*types.Group{ + {ID: "dns-grpA", Name: "DNS-A", Peers: []string{peer1.ID}}, + {ID: "dns-grpB", Name: "DNS-B", Peers: []string{peer2.ID}}, + } { + err := manager.CreateGroup(ctx, accountID, userID, g) + require.NoError(t, err) + } + + updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) + updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) + updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + // Save DNS settings that only affects grpA. peer2 and peer3 should not be affected. + t.Run("dns settings change only affects linked peers", func(t *testing.T) { + done := make(chan struct{}) + go func() { + peerShouldReceiveUpdate(t, updMsg1) + peerShouldNotReceiveUpdate(t, updMsg2) + peerShouldNotReceiveUpdate(t, updMsg3) + close(done) + }() + + err := manager.SaveDNSSettings(ctx, accountID, userID, &types.DNSSettings{ + DisabledManagementGroups: []string{"dns-grpA"}, + }) + assert.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout") + } + }) +} + +// TestAffectedPeers_UnlinkedGroupChange_NoUpdateIntegration tests the full integration: +// updating a group that is NOT referenced by any policy/route/ns/dns should not send +// updates to any peer. +func TestAffectedPeers_UnlinkedGroupChange_NoUpdateIntegration(t *testing.T) { + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + err := manager.Store.DeletePolicy(ctx, accountID, p.ID) + require.NoError(t, err) + } + + err = manager.CreateGroup(ctx, accountID, userID, &types.Group{ + ID: "unlinked-grp", + Name: "Unlinked", + Peers: []string{peer1.ID}, + }) + require.NoError(t, err) + + updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) + updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) + updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + t.Run("updating unlinked group sends no peer updates", func(t *testing.T) { + done := make(chan struct{}) + go func() { + peerShouldNotReceiveUpdate(t, updMsg1) + peerShouldNotReceiveUpdate(t, updMsg2) + peerShouldNotReceiveUpdate(t, updMsg3) + close(done) + }() + + err := manager.UpdateGroup(ctx, accountID, userID, &types.Group{ + ID: "unlinked-grp", + Name: "Unlinked", + Peers: []string{peer1.ID, peer2.ID}, + }) + assert.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout") + } + }) +} + +// TestAffectedPeers_NetworkRouter_UnrelatedPeerNoUpdate verifies that when a network +// router is added with specific peer groups, only peers in those groups (and policy +// sources for resources) get updates. Unrelated peers should not. +func TestAffectedPeers_NetworkRouter_UnrelatedPeerNoUpdate(t *testing.T) { + // Use custom setup: delete default policy BEFORE adding peers so that + // AddPeer's BufferUpdateAffectedPeers finds no affected peers and + // doesn't schedule async updates that race with the test. + manager, updateManager, err := createManager(t) + require.NoError(t, err) + + ctx := context.Background() + + account, err := createAccount(manager, "nr_test_account", userID, "") + require.NoError(t, err) + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + err := manager.Store.DeletePolicy(ctx, accountID, p.ID) + require.NoError(t, err) + } + + setupKey, err := manager.CreateSetupKey(ctx, accountID, "test-key", types.SetupKeyReusable, time.Hour, nil, 999, userID, false, false) + require.NoError(t, err) + + peer1 := addPeerToAccount(t, manager, accountID, setupKey.Key) + peer2 := addPeerToAccount(t, manager, accountID, setupKey.Key) + peer3 := addPeerToAccount(t, manager, accountID, setupKey.Key) + + for _, g := range []*types.Group{ + {ID: "nr-grpA", Name: "NR-A", Peers: []string{peer1.ID}}, + {ID: "nr-grpB", Name: "NR-B", Peers: []string{peer2.ID}}, + } { + err := manager.CreateGroup(ctx, accountID, userID, g) + require.NoError(t, err) + } + + net1 := &networkTypes.Network{ + ID: "nr-net-test", + AccountID: accountID, + Name: "nr-test-network", + } + err = manager.Store.SaveNetwork(ctx, net1) + require.NoError(t, err) + + err = manager.Store.SaveNetworkRouter(ctx, &routerTypes.NetworkRouter{ + ID: "nr-router-test", + NetworkID: net1.ID, + AccountID: accountID, + PeerGroups: []string{"nr-grpA"}, + }) + require.NoError(t, err) + + updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) + updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) + updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + // When the group linked to the network router changes, only peers in that + // group should be updated. Peer2 is unrelated. Peer3 is added to the + // router's group so it should also receive an update. + t.Run("network router group change only affects linked peers", func(t *testing.T) { + done := make(chan struct{}) + go func() { + peerShouldReceiveUpdate(t, updMsg1) + peerShouldNotReceiveUpdate(t, updMsg2) + peerShouldReceiveUpdate(t, updMsg3) + close(done) + }() + + // Updating the group linked to router should affect peer1 and peer3 (now in nr-grpA). + err = manager.UpdateGroup(ctx, accountID, userID, &types.Group{ + ID: "nr-grpA", + Name: "NR-A", + Peers: []string{peer1.ID, peer3.ID}, + }) + assert.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout") + } + }) +} + +// TestAffectedPeers_MultipleIsolatedEntities_OnlyLinkedPeersUpdated creates multiple +// isolated entities (policy for peer1<->peer2, route for peer3) and verifies that +// changing one entity's groups only affects its peers. +func TestAffectedPeers_MultipleIsolatedEntities_OnlyLinkedPeersUpdated(t *testing.T) { + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + err := manager.Store.DeletePolicy(ctx, accountID, p.ID) + require.NoError(t, err) + } + + for _, g := range []*types.Group{ + {ID: "iso-grpA", Name: "ISO-A", Peers: []string{peer1.ID}}, + {ID: "iso-grpB", Name: "ISO-B", Peers: []string{peer2.ID}}, + {ID: "iso-grpC", Name: "ISO-C", Peers: []string{peer3.ID}}, + } { + err := manager.CreateGroup(ctx, accountID, userID, g) + require.NoError(t, err) + } + + // Policy: peer1 <-> peer2 + _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, Rules: []*types.PolicyRule{ { - Sources: []string{"g1", "g2"}, - Destinations: []string{"g3"}, + Enabled: true, + Sources: []string{"iso-grpA"}, + Destinations: []string{"iso-grpB"}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, }, }, - } + }, true) + require.NoError(t, err) - tests := []struct { - name string - groupSet map[string]struct{} - want bool - }{ - {"matches source", map[string]struct{}{"g1": {}}, true}, - {"matches destination", map[string]struct{}{"g3": {}}, true}, - {"no match", map[string]struct{}{"g4": {}}, false}, - {"empty set", map[string]struct{}{}, false}, - } + // Route: only peer3's group as distribution group + _, err = manager.CreateRoute(ctx, accountID, + netip.MustParsePrefix("10.20.0.0/24"), + route.IPv4Network, + nil, + "", + []string{"iso-grpC"}, + "isolated route", + "isonet2", + false, + 9999, + []string{"iso-grpC"}, + nil, + true, + userID, + false, + false, + ) + require.NoError(t, err) - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := policyReferencesGroups(policy, tt.groupSet) - assert.Equal(t, tt.want, got) + updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) + updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) + updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + // Updating policy group (iso-grpA) should affect peer1+peer2 but NOT peer3 + t.Run("policy group change does not affect route-only peer", func(t *testing.T) { + done := make(chan struct{}) + go func() { + peerShouldReceiveUpdate(t, updMsg1) + peerShouldReceiveUpdate(t, updMsg2) + peerShouldNotReceiveUpdate(t, updMsg3) + close(done) + }() + + err := manager.UpdateGroup(ctx, accountID, userID, &types.Group{ + ID: "iso-grpA", + Name: "ISO-A-updated", + Peers: []string{peer1.ID}, }) - } + assert.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout") + } + }) } -func TestRouteReferencesGroups(t *testing.T) { - r := &route.Route{ - Groups: []string{"g1"}, - PeerGroups: []string{"g2"}, - AccessControlGroups: []string{"g3"}, +// TestAffectedPeers_DeleteRoute_UnrelatedPeerNoUpdate verifies that deleting a route +// only sends updates to peers in the route's groups. +func TestAffectedPeers_DeleteRoute_UnrelatedPeerNoUpdate(t *testing.T) { + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + err := manager.Store.DeletePolicy(ctx, accountID, p.ID) + require.NoError(t, err) } - tests := []struct { - name string - groupSet map[string]struct{} - want bool - }{ - {"matches groups", map[string]struct{}{"g1": {}}, true}, - {"matches peerGroups", map[string]struct{}{"g2": {}}, true}, - {"matches accessControl", map[string]struct{}{"g3": {}}, true}, - {"no match", map[string]struct{}{"g4": {}}, false}, + for _, g := range []*types.Group{ + {ID: "del-rt-grpA", Name: "Del-Rt-A", Peers: []string{peer1.ID}}, + {ID: "del-rt-grpB", Name: "Del-Rt-B", Peers: []string{peer2.ID}}, + } { + err := manager.CreateGroup(ctx, accountID, userID, g) + require.NoError(t, err) } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := routeReferencesGroups(r, tt.groupSet) - assert.Equal(t, tt.want, got) - }) - } + newRoute, err := manager.CreateRoute(ctx, accountID, + netip.MustParsePrefix("10.30.0.0/24"), + route.IPv4Network, + nil, + "", + []string{"del-rt-grpA"}, + "deletable route", + "delnet", + false, + 9999, + []string{"del-rt-grpB"}, + nil, + true, + userID, + false, + false, + ) + require.NoError(t, err) + + updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) + updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) + updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + t.Run("delete route only affects linked peers", func(t *testing.T) { + done := make(chan struct{}) + go func() { + peerShouldReceiveUpdate(t, updMsg1) + peerShouldReceiveUpdate(t, updMsg2) + peerShouldNotReceiveUpdate(t, updMsg3) + close(done) + }() + + err := manager.DeleteRoute(ctx, accountID, newRoute.ID, userID) + assert.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout") + } + }) } -func TestRouterReferencesGroups(t *testing.T) { - router := &routerTypes.NetworkRouter{ - PeerGroups: []string{"g1", "g2"}, +// TestAffectedPeers_DeletePolicy_UnrelatedPeerNoUpdate verifies that deleting a policy +// only sends updates to peers in the policy's groups. +func TestAffectedPeers_DeletePolicy_UnrelatedPeerNoUpdate(t *testing.T) { + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + err := manager.Store.DeletePolicy(ctx, accountID, p.ID) + require.NoError(t, err) } - tests := []struct { - name string - groupSet map[string]struct{} - want bool - }{ - {"matches", map[string]struct{}{"g1": {}}, true}, - {"no match", map[string]struct{}{"g3": {}}, false}, + for _, g := range []*types.Group{ + {ID: "del-pol-grpA", Name: "Del-Pol-A", Peers: []string{peer1.ID}}, + {ID: "del-pol-grpB", Name: "Del-Pol-B", Peers: []string{peer2.ID}}, + } { + err := manager.CreateGroup(ctx, accountID, userID, g) + require.NoError(t, err) } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := routerReferencesGroups(router, tt.groupSet) - assert.Equal(t, tt.want, got) - }) - } + policy, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{"del-pol-grpA"}, + Destinations: []string{"del-pol-grpB"}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) + updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) + updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + t.Run("delete policy only affects linked peers", func(t *testing.T) { + done := make(chan struct{}) + go func() { + peerShouldReceiveUpdate(t, updMsg1) + peerShouldReceiveUpdate(t, updMsg2) + peerShouldNotReceiveUpdate(t, updMsg3) + close(done) + }() + + err := manager.DeletePolicy(ctx, accountID, policy.ID, userID) + assert.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout") + } + }) } -func addPeerToAccount(t *testing.T, manager *DefaultAccountManager, accountID, setupKeyKey string) *nbpeer.Peer { +// TestAffectedPeers_DeleteNameServer_UnrelatedPeerNoUpdate verifies that deleting a +// nameserver group only sends updates to peers in its groups. +func TestAffectedPeers_DeleteNameServer_UnrelatedPeerNoUpdate(t *testing.T) { + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + err := manager.Store.DeletePolicy(ctx, accountID, p.ID) + require.NoError(t, err) + } + + err = manager.CreateGroup(ctx, accountID, userID, &types.Group{ + ID: "del-ns-grpA", + Name: "Del-NS-A", + Peers: []string{peer1.ID}, + }) + require.NoError(t, err) + + nsGroup, err := manager.CreateNameServerGroup(ctx, accountID, "del-ns", "Del NS", + []nbdns.NameServer{{ + IP: netip.MustParseAddr("8.8.4.4"), + NSType: nbdns.UDPNameServerType, + Port: nbdns.DefaultDNSPort, + }}, + []string{"del-ns-grpA"}, + true, nil, true, userID, false, + ) + require.NoError(t, err) + + updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) + updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) + updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + t.Run("delete nameserver group only affects linked peers", func(t *testing.T) { + done := make(chan struct{}) + go func() { + peerShouldReceiveUpdate(t, updMsg1) + peerShouldNotReceiveUpdate(t, updMsg2) + peerShouldNotReceiveUpdate(t, updMsg3) + close(done) + }() + + err := manager.DeleteNameServerGroup(ctx, accountID, nsGroup.ID, userID) + assert.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout") + } + }) +} + +func addPeerToAccount(t *testing.T, manager *DefaultAccountManager, _, setupKeyKey string) *nbpeer.Peer { t.Helper() key, err := wgtypes.GeneratePrivateKey() From 46494bd8604b9fe5f7617b4f26c7a0837f4f69a7 Mon Sep 17 00:00:00 2001 From: pascal Date: Thu, 7 May 2026 16:08:45 +0200 Subject: [PATCH 08/28] bugfixes --- .../controllers/network_map/controller/controller.go | 2 +- management/server/networks/routers/manager.go | 5 ++++- management/server/posture_checks.go | 2 ++ 3 files changed, 7 insertions(+), 2 deletions(-) diff --git a/management/internals/controllers/network_map/controller/controller.go b/management/internals/controllers/network_map/controller/controller.go index 150c59d4b..8ba7c3a49 100644 --- a/management/internals/controllers/network_map/controller/controller.go +++ b/management/internals/controllers/network_map/controller/controller.go @@ -204,7 +204,7 @@ func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID strin c.metrics.CountCalcPeerNetworkMapDuration(time.Since(start)) - proxyNetworkMap, ok := proxyNetworkMaps[peer.ID] + proxyNetworkMap, ok := proxyNetworkMaps[p.ID] if ok { remotePeerNetworkMap.Merge(proxyNetworkMap) } diff --git a/management/server/networks/routers/manager.go b/management/server/networks/routers/manager.go index 7ee9dc281..6b33f558e 100644 --- a/management/server/networks/routers/manager.go +++ b/management/server/networks/routers/manager.go @@ -177,7 +177,10 @@ func (m *managerImpl) UpdateRouter(ctx context.Context, userID string, router *t } allPeerGroups := router.PeerGroups - directPeers := []string{router.Peer} + var directPeers []string + if router.Peer != "" { + directPeers = append(directPeers, router.Peer) + } oldRouter, err := transaction.GetNetworkRouterByID(ctx, store.LockingStrengthNone, router.AccountID, router.ID) if err == nil { allPeerGroups = append(allPeerGroups, oldRouter.PeerGroups...) diff --git a/management/server/posture_checks.go b/management/server/posture_checks.go index cc6abc943..f22bc8d14 100644 --- a/management/server/posture_checks.go +++ b/management/server/posture_checks.go @@ -5,6 +5,7 @@ import ( "slices" "github.com/rs/xid" + log "github.com/sirupsen/logrus" "github.com/netbirdio/netbird/management/server/activity" "github.com/netbirdio/netbird/management/server/permissions/modules" @@ -134,6 +135,7 @@ func (am *DefaultAccountManager) ListPostureChecks(ctx context.Context, accountI func collectPostureCheckAffectedGroupsAndPeers(ctx context.Context, transaction store.Store, accountID, postureCheckID string) (groupIDs []string, directPeerIDs []string) { policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) if err != nil { + log.WithContext(ctx).Errorf("failed to get policies for posture check affected peers resolution: %v", err) return nil, nil } From 550ae5558ebbe8f83143862f9f1155e3fd9d96a5 Mon Sep 17 00:00:00 2001 From: pascal Date: Thu, 7 May 2026 16:24:54 +0200 Subject: [PATCH 09/28] update after merge --- .../controllers/network_map/controller/controller.go | 6 +++++- management/internals/controllers/network_map/interface.go | 2 +- .../internals/controllers/network_map/interface_mock.go | 8 ++++---- management/server/account.go | 4 +++- management/server/account/manager.go | 2 +- management/server/account/manager_mock.go | 8 ++++---- management/server/mock_server/account_mock.go | 6 +++--- management/server/peer.go | 4 ++-- 8 files changed, 23 insertions(+), 17 deletions(-) diff --git a/management/internals/controllers/network_map/controller/controller.go b/management/internals/controllers/network_map/controller/controller.go index 8ba7c3a49..9ee33692b 100644 --- a/management/internals/controllers/network_map/controller/controller.go +++ b/management/internals/controllers/network_map/controller/controller.go @@ -512,11 +512,15 @@ func (c *Controller) BufferUpdateAccountPeers(ctx context.Context, accountID str } // BufferUpdateAffectedPeers accumulates peer IDs and flushes them after the buffer interval. -func (c *Controller) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error { +func (c *Controller) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string, reason types.UpdateReason) error { if len(peerIDs) == 0 { return nil } + if c.accountManagerMetrics != nil { + c.accountManagerMetrics.CountUpdateAccountPeersTriggered(string(reason.Resource), string(reason.Operation)) + } + log.WithContext(ctx).Tracef("buffer updating %d affected peers for account %s from %s", len(peerIDs), accountID, util.GetCallerName()) bufUpd, _ := c.affectedPeerUpdateLocks.LoadOrStore(accountID, &bufferAffectedUpdate{ diff --git a/management/internals/controllers/network_map/interface.go b/management/internals/controllers/network_map/interface.go index 95bed7533..dbdd87708 100644 --- a/management/internals/controllers/network_map/interface.go +++ b/management/internals/controllers/network_map/interface.go @@ -20,7 +20,7 @@ const ( type Controller interface { UpdateAccountPeers(ctx context.Context, accountID string, reason types.UpdateReason) error UpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error - BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error + BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string, reason types.UpdateReason) error UpdateAccountPeer(ctx context.Context, accountId string, peerId string) error BufferUpdateAccountPeers(ctx context.Context, accountID string, reason types.UpdateReason) error GetValidatedPeerWithMap(ctx context.Context, isRequiresApproval bool, accountID string, p *nbpeer.Peer) (*nbpeer.Peer, *types.NetworkMap, []*posture.Checks, int64, error) diff --git a/management/internals/controllers/network_map/interface_mock.go b/management/internals/controllers/network_map/interface_mock.go index 560a4b42b..a67156719 100644 --- a/management/internals/controllers/network_map/interface_mock.go +++ b/management/internals/controllers/network_map/interface_mock.go @@ -58,17 +58,17 @@ func (mr *MockControllerMockRecorder) BufferUpdateAccountPeers(ctx, accountID, r } // BufferUpdateAffectedPeers mocks base method. -func (m *MockController) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error { +func (m *MockController) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string, reason types.UpdateReason) error { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "BufferUpdateAffectedPeers", ctx, accountID, peerIDs) + ret := m.ctrl.Call(m, "BufferUpdateAffectedPeers", ctx, accountID, peerIDs, reason) ret0, _ := ret[0].(error) return ret0 } // BufferUpdateAffectedPeers indicates an expected call of BufferUpdateAffectedPeers. -func (mr *MockControllerMockRecorder) BufferUpdateAffectedPeers(ctx, accountID, peerIDs any) *gomock.Call { +func (mr *MockControllerMockRecorder) BufferUpdateAffectedPeers(ctx, accountID, peerIDs, reason any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "BufferUpdateAffectedPeers", reflect.TypeOf((*MockController)(nil).BufferUpdateAffectedPeers), ctx, accountID, peerIDs) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "BufferUpdateAffectedPeers", reflect.TypeOf((*MockController)(nil).BufferUpdateAffectedPeers), ctx, accountID, peerIDs, reason) } // CountStreams mocks base method. diff --git a/management/server/account.go b/management/server/account.go index b2d62bc24..000a043c8 100644 --- a/management/server/account.go +++ b/management/server/account.go @@ -2589,7 +2589,9 @@ func (am *DefaultAccountManager) UpdatePeerIPv6(ctx context.Context, accountID, } if updateNetworkMap { - if err := am.networkMapController.OnPeersUpdated(ctx, accountID, []string{peerID}); err != nil { + changedPeerIDs := []string{peerID} + affectedPeerIDs := am.resolveAffectedPeersForPeerChanges(ctx, am.Store, accountID, changedPeerIDs) + if err := am.networkMapController.OnPeersUpdated(ctx, accountID, changedPeerIDs, affectedPeerIDs); err != nil { return fmt.Errorf("notify network map controller: %w", err) } } diff --git a/management/server/account/manager.go b/management/server/account/manager.go index 123dd1829..7832eb652 100644 --- a/management/server/account/manager.go +++ b/management/server/account/manager.go @@ -127,7 +127,7 @@ type Manager interface { DeleteSetupKey(ctx context.Context, accountID, userID, keyID string) error UpdateAccountPeers(ctx context.Context, accountID string, reason types.UpdateReason) UpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) - BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) + BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string, reason types.UpdateReason) BufferUpdateAccountPeers(ctx context.Context, accountID string, reason types.UpdateReason) BuildUserInfosForAccount(ctx context.Context, accountID, initiatorUserID string, accountUsers []*types.User) (map[string]*types.UserInfo, error) SyncUserJWTGroups(ctx context.Context, userAuth auth.UserAuth) error diff --git a/management/server/account/manager_mock.go b/management/server/account/manager_mock.go index 3668bd5df..c11252db1 100644 --- a/management/server/account/manager_mock.go +++ b/management/server/account/manager_mock.go @@ -123,15 +123,15 @@ func (mr *MockManagerMockRecorder) BufferUpdateAccountPeers(ctx, accountID, reas } // BufferUpdateAffectedPeers mocks base method. -func (m *MockManager) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) { +func (m *MockManager) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string, reason types.UpdateReason) { m.ctrl.T.Helper() - m.ctrl.Call(m, "BufferUpdateAffectedPeers", ctx, accountID, peerIDs) + m.ctrl.Call(m, "BufferUpdateAffectedPeers", ctx, accountID, peerIDs, reason) } // BufferUpdateAffectedPeers indicates an expected call of BufferUpdateAffectedPeers. -func (mr *MockManagerMockRecorder) BufferUpdateAffectedPeers(ctx, accountID, peerIDs interface{}) *gomock.Call { +func (mr *MockManagerMockRecorder) BufferUpdateAffectedPeers(ctx, accountID, peerIDs, reason interface{}) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "BufferUpdateAffectedPeers", reflect.TypeOf((*MockManager)(nil).BufferUpdateAffectedPeers), ctx, accountID, peerIDs) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "BufferUpdateAffectedPeers", reflect.TypeOf((*MockManager)(nil).BufferUpdateAffectedPeers), ctx, accountID, peerIDs, reason) } // BuildUserInfosForAccount mocks base method. diff --git a/management/server/mock_server/account_mock.go b/management/server/mock_server/account_mock.go index bbde919a8..8e7ce0f51 100644 --- a/management/server/mock_server/account_mock.go +++ b/management/server/mock_server/account_mock.go @@ -131,7 +131,7 @@ type MockAccountManager struct { AllowSyncFunc func(string, uint64) bool UpdateAccountPeersFunc func(ctx context.Context, accountID string, reason types.UpdateReason) UpdateAffectedPeersFunc func(ctx context.Context, accountID string, peerIDs []string) - BufferUpdateAffectedPeersFunc func(ctx context.Context, accountID string, peerIDs []string) + BufferUpdateAffectedPeersFunc func(ctx context.Context, accountID string, peerIDs []string, reason types.UpdateReason) BufferUpdateAccountPeersFunc func(ctx context.Context, accountID string, reason types.UpdateReason) RecalculateNetworkMapCacheFunc func(ctx context.Context, accountId string) error @@ -215,9 +215,9 @@ func (am *MockAccountManager) UpdateAffectedPeers(ctx context.Context, accountID } } -func (am *MockAccountManager) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) { +func (am *MockAccountManager) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string, reason types.UpdateReason) { if am.BufferUpdateAffectedPeersFunc != nil { - am.BufferUpdateAffectedPeersFunc(ctx, accountID, peerIDs) + am.BufferUpdateAffectedPeersFunc(ctx, accountID, peerIDs, reason) } } diff --git a/management/server/peer.go b/management/server/peer.go index 65d433a88..08d5b6deb 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -1328,8 +1328,8 @@ func (am *DefaultAccountManager) resolvePeerIDs(ctx context.Context, s store.Sto } // BufferUpdateAffectedPeers accumulates peer IDs and flushes them after the buffer interval. -func (am *DefaultAccountManager) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string) { - _ = am.networkMapController.BufferUpdateAffectedPeers(ctx, accountID, peerIDs) +func (am *DefaultAccountManager) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string, reason types.UpdateReason) { + _ = am.networkMapController.BufferUpdateAffectedPeers(ctx, accountID, peerIDs, reason) } // resolveAffectedPeersForPeerChanges resolves changed peer IDs into the full set of affected peer IDs. From ec476d50720f2651d0a848359806da181d196558 Mon Sep 17 00:00:00 2001 From: pascal Date: Thu, 7 May 2026 16:55:45 +0200 Subject: [PATCH 10/28] extend logging --- .../network_map/controller/controller.go | 6 ++++- management/server/affected_peers_test.go | 22 ++++++++--------- management/server/dns.go | 3 +++ management/server/group.go | 24 +++++++++++++++++++ management/server/nameserver.go | 10 ++++++++ management/server/networks/manager.go | 7 ++++++ .../server/networks/resources/manager.go | 12 ++++++++++ management/server/networks/routers/manager.go | 12 ++++++++++ management/server/peer.go | 4 ++-- management/server/policy.go | 23 +++++++++++++----- management/server/posture_checks.go | 7 +++++- management/server/route.go | 21 ++++++++++++---- 12 files changed, 126 insertions(+), 25 deletions(-) diff --git a/management/internals/controllers/network_map/controller/controller.go b/management/internals/controllers/network_map/controller/controller.go index 9ee33692b..b17036c6b 100644 --- a/management/internals/controllers/network_map/controller/controller.go +++ b/management/internals/controllers/network_map/controller/controller.go @@ -285,7 +285,7 @@ func (c *Controller) UpdateAffectedPeers(ctx context.Context, accountID string, } func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error { - log.WithContext(ctx).Tracef("updating %d affected peers for account %s from %s", len(peerIDs), accountID, util.GetCallerName()) + log.WithContext(ctx).Tracef("sendUpdateForAffectedPeers: account %s, %d affected peers: %v (caller: %s)", accountID, len(peerIDs), peerIDs, util.GetCallerName()) affected := make(map[string]struct{}, len(peerIDs)) for _, id := range peerIDs { @@ -300,6 +300,7 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s } } if !hasConnected { + log.WithContext(ctx).Tracef("sendUpdateForAffectedPeers: no connected peers among %v, skipping", peerIDs) return nil } @@ -318,9 +319,12 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s } if len(peersToUpdate) == 0 { + log.WithContext(ctx).Tracef("sendUpdateForAffectedPeers: no peers to update (affected peers not found in account or no channels)") return nil } + log.WithContext(ctx).Tracef("sendUpdateForAffectedPeers: sending network map to %d connected peers", len(peersToUpdate)) + approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra) if err != nil { return fmt.Errorf("failed to get validate peers: %v", err) diff --git a/management/server/affected_peers_test.go b/management/server/affected_peers_test.go index 0004fe5c1..13bb990c8 100644 --- a/management/server/affected_peers_test.go +++ b/management/server/affected_peers_test.go @@ -437,7 +437,7 @@ func TestCollectPolicyAffectedGroups_Basic(t *testing.T) { }, }, } - groups, directPeers := collectPolicyAffectedGroupsAndPeers(policy) + groups, directPeers := collectPolicyAffectedGroupsAndPeers(context.Background(), policy) assert.ElementsMatch(t, []string{"g1", "g2", "g3"}, groups) assert.Empty(t, directPeers) } @@ -453,13 +453,13 @@ func TestCollectPolicyAffectedGroups_WithPeerResources(t *testing.T) { }, }, } - groups, directPeers := collectPolicyAffectedGroupsAndPeers(policy) + groups, directPeers := collectPolicyAffectedGroupsAndPeers(context.Background(), policy) assert.ElementsMatch(t, []string{"g1", "g2"}, groups) assert.ElementsMatch(t, []string{"p1", "p2"}, directPeers) } func TestCollectPolicyAffectedGroups_NilPolicy(t *testing.T) { - groups, directPeers := collectPolicyAffectedGroupsAndPeers(nil) + groups, directPeers := collectPolicyAffectedGroupsAndPeers(context.Background(), nil) assert.Nil(t, groups) assert.Nil(t, directPeers) } @@ -471,7 +471,7 @@ func TestCollectPolicyAffectedGroups_MultipleRules(t *testing.T) { {Sources: []string{"g3"}, Destinations: []string{"g4"}}, }, } - groups, _ := collectPolicyAffectedGroupsAndPeers(policy) + groups, _ := collectPolicyAffectedGroupsAndPeers(context.Background(), policy) assert.ElementsMatch(t, []string{"g1", "g2", "g3", "g4"}, groups) } @@ -486,13 +486,13 @@ func TestCollectPolicyAffectedGroups_MultiplePolicies(t *testing.T) { {Sources: []string{"g3"}, Destinations: []string{"g4"}}, }, } - groups, _ := collectPolicyAffectedGroupsAndPeers(new, old) + groups, _ := collectPolicyAffectedGroupsAndPeers(context.Background(), new, old) assert.ElementsMatch(t, []string{"g1", "g2", "g3", "g4"}, groups) } func TestCollectPolicyAffectedGroups_EmptyRules(t *testing.T) { policy := &types.Policy{Rules: []*types.PolicyRule{}} - groups, directPeers := collectPolicyAffectedGroupsAndPeers(policy) + groups, directPeers := collectPolicyAffectedGroupsAndPeers(context.Background(), policy) assert.Empty(t, groups) assert.Empty(t, directPeers) } @@ -507,7 +507,7 @@ func TestCollectPolicyAffectedGroups_NonPeerResource(t *testing.T) { }, }, } - groups, directPeers := collectPolicyAffectedGroupsAndPeers(policy) + groups, directPeers := collectPolicyAffectedGroupsAndPeers(context.Background(), policy) assert.ElementsMatch(t, []string{"g1", "g2"}, groups) assert.Empty(t, directPeers, "domain resource type should not produce direct peer IDs") } @@ -522,7 +522,7 @@ func TestCollectRouteAffectedGroups_Basic(t *testing.T) { PeerGroups: []string{"g2"}, AccessControlGroups: []string{"g3"}, } - groups, directPeers := collectRouteAffectedGroupsAndPeers(r) + groups, directPeers := collectRouteAffectedGroupsAndPeers(context.Background(), r) assert.ElementsMatch(t, []string{"g1", "g2", "g3"}, groups) assert.Empty(t, directPeers) } @@ -532,13 +532,13 @@ func TestCollectRouteAffectedGroups_WithDirectPeer(t *testing.T) { Groups: []string{"g1"}, Peer: "p1", } - groups, directPeers := collectRouteAffectedGroupsAndPeers(r) + groups, directPeers := collectRouteAffectedGroupsAndPeers(context.Background(), r) assert.ElementsMatch(t, []string{"g1"}, groups) assert.ElementsMatch(t, []string{"p1"}, directPeers) } func TestCollectRouteAffectedGroups_NilRoute(t *testing.T) { - groups, directPeers := collectRouteAffectedGroupsAndPeers(nil) + groups, directPeers := collectRouteAffectedGroupsAndPeers(context.Background(), nil) assert.Nil(t, groups) assert.Nil(t, directPeers) } @@ -552,7 +552,7 @@ func TestCollectRouteAffectedGroups_MultipleRoutes(t *testing.T) { Groups: []string{"g2"}, PeerGroups: []string{"g3"}, } - groups, directPeers := collectRouteAffectedGroupsAndPeers(new, old) + groups, directPeers := collectRouteAffectedGroupsAndPeers(context.Background(), new, old) assert.ElementsMatch(t, []string{"g1", "g2", "g3"}, groups) assert.ElementsMatch(t, []string{"p1"}, directPeers) } diff --git a/management/server/dns.go b/management/server/dns.go index 1e213ffbb..144aa5ca5 100644 --- a/management/server/dns.go +++ b/management/server/dns.go @@ -84,7 +84,10 @@ func (am *DefaultAccountManager) SaveDNSSettings(ctx context.Context, accountID } if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("SaveDNSSettings: updating %d affected peers: %v", len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("SaveDNSSettings: no affected peers") } return nil diff --git a/management/server/group.go b/management/server/group.go index df1b60156..5d15eed39 100644 --- a/management/server/group.go +++ b/management/server/group.go @@ -115,7 +115,10 @@ func (am *DefaultAccountManager) CreateGroup(ctx context.Context, accountID, use } if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("CreateGroup %s: updating %d affected peers: %v", newGroup.ID, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("CreateGroup %s: no affected peers", newGroup.ID) } return nil @@ -185,7 +188,10 @@ func (am *DefaultAccountManager) UpdateGroup(ctx context.Context, accountID, use } if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("UpdateGroup %s: updating %d affected peers: %v", newGroup.ID, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("UpdateGroup %s: no affected peers", newGroup.ID) } return nil @@ -249,7 +255,10 @@ func (am *DefaultAccountManager) CreateGroups(ctx context.Context, accountID, us allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, am.Store, accountID, groupIDs) affectedPeerIDs := am.resolvePeerIDs(ctx, am.Store, accountID, allGroupIDs, directPeerIDs) if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("CreateGroups %v: updating %d affected peers: %v", groupIDs, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("CreateGroups %v: no affected peers", groupIDs) } return globalErr @@ -293,7 +302,10 @@ func (am *DefaultAccountManager) UpdateGroups(ctx context.Context, accountID, us allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, am.Store, accountID, groupIDs) affectedPeerIDs := am.resolvePeerIDs(ctx, am.Store, accountID, allGroupIDs, directPeerIDs) if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("UpdateGroups %v: updating %d affected peers: %v", groupIDs, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("UpdateGroups %v: no affected peers", groupIDs) } return globalErr @@ -498,7 +510,10 @@ func (am *DefaultAccountManager) GroupAddPeer(ctx context.Context, accountID, gr } if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("GroupAddPeer group=%s peer=%s: updating %d affected peers: %v", groupID, peerID, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("GroupAddPeer group=%s peer=%s: no affected peers", groupID, peerID) } return nil @@ -534,7 +549,10 @@ func (am *DefaultAccountManager) GroupAddResource(ctx context.Context, accountID } if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("GroupAddResource group=%s resource=%s: updating %d affected peers: %v", groupID, resource.ID, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("GroupAddResource group=%s resource=%s: no affected peers", groupID, resource.ID) } return nil @@ -565,7 +583,10 @@ func (am *DefaultAccountManager) GroupDeletePeer(ctx context.Context, accountID, } if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("GroupDeletePeer group=%s peer=%s: updating %d affected peers: %v", groupID, peerID, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("GroupDeletePeer group=%s peer=%s: no affected peers", groupID, peerID) } return nil @@ -601,7 +622,10 @@ func (am *DefaultAccountManager) GroupDeleteResource(ctx context.Context, accoun } if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("GroupDeleteResource group=%s resource=%s: updating %d affected peers: %v", groupID, resource.ID, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("GroupDeleteResource group=%s resource=%s: no affected peers", groupID, resource.ID) } return nil diff --git a/management/server/nameserver.go b/management/server/nameserver.go index 823fc72d5..8ba1c2afc 100644 --- a/management/server/nameserver.go +++ b/management/server/nameserver.go @@ -9,6 +9,7 @@ import ( "unicode/utf8" "github.com/rs/xid" + log "github.com/sirupsen/logrus" nbdns "github.com/netbirdio/netbird/dns" "github.com/netbirdio/netbird/management/server/activity" @@ -80,7 +81,10 @@ func (am *DefaultAccountManager) CreateNameServerGroup(ctx context.Context, acco am.StoreEvent(ctx, userID, newNSGroup.ID, accountID, activity.NameserverGroupCreated, newNSGroup.EventMeta()) if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("CreateNameServerGroup %s: updating %d affected peers: %v", newNSGroup.ID, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("CreateNameServerGroup %s: no affected peers", newNSGroup.ID) } return newNSGroup.Copy(), nil @@ -129,7 +133,10 @@ func (am *DefaultAccountManager) SaveNameServerGroup(ctx context.Context, accoun am.StoreEvent(ctx, userID, nsGroupToSave.ID, accountID, activity.NameserverGroupUpdated, nsGroupToSave.EventMeta()) if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("SaveNameServerGroup %s: updating %d affected peers: %v", nsGroupToSave.ID, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("SaveNameServerGroup %s: no affected peers", nsGroupToSave.ID) } return nil @@ -169,7 +176,10 @@ func (am *DefaultAccountManager) DeleteNameServerGroup(ctx context.Context, acco am.StoreEvent(ctx, userID, nsGroup.ID, accountID, activity.NameserverGroupDeleted, nsGroup.EventMeta()) if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("DeleteNameServerGroup %s: updating %d affected peers: %v", nsGroupID, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("DeleteNameServerGroup %s: no affected peers", nsGroupID) } return nil diff --git a/management/server/networks/manager.go b/management/server/networks/manager.go index 17ea0ddaa..0f21ea9ba 100644 --- a/management/server/networks/manager.go +++ b/management/server/networks/manager.go @@ -224,7 +224,10 @@ func (m *managerImpl) DeleteNetwork(ctx context.Context, accountID, userID, netw if affectedData != nil { affectedPeerIDs := resolveNetworkAffectedPeers(ctx, m.store, accountID, affectedData) if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("DeleteNetwork %s: updating %d affected peers: %v", networkID, len(affectedPeerIDs), affectedPeerIDs) go m.accountManager.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("DeleteNetwork %s: no affected peers", networkID) } } @@ -233,6 +236,8 @@ func (m *managerImpl) DeleteNetwork(ctx context.Context, accountID, userID, netw // resolveNetworkAffectedPeers computes affected peer IDs from preloaded data outside the transaction. func resolveNetworkAffectedPeers(ctx context.Context, s store.Store, accountID string, data *networkAffectedPeersData) []string { + log.WithContext(ctx).Tracef("resolveNetworkAffectedPeers: routerPeerGroups=%v, resourceGroupIDs=%v, directPeerIDs=%v, policies=%d", + data.routerPeerGroups, data.resourceGroupIDs, data.directPeerIDs, len(data.policies)) groupSet := make(map[string]struct{}) for _, gID := range data.routerPeerGroups { @@ -275,6 +280,7 @@ func resolveNetworkAffectedPeers(ctx context.Context, s store.Store, accountID s groupIDs = append(groupIDs, gID) } + log.WithContext(ctx).Tracef("resolveNetworkAffectedPeers: resolved groupIDs=%v", groupIDs) peerIDs, err := s.GetPeerIDsByGroups(ctx, accountID, groupIDs) if err != nil { log.WithContext(ctx).Errorf("failed to resolve peer IDs: %v", err) @@ -294,6 +300,7 @@ func resolveNetworkAffectedPeers(ctx context.Context, s store.Store, accountID s } } + log.WithContext(ctx).Tracef("resolveNetworkAffectedPeers: result %d peers: %v", len(peerIDs), peerIDs) return peerIDs } diff --git a/management/server/networks/resources/manager.go b/management/server/networks/resources/manager.go index 4d016cdf6..23f4f98a4 100644 --- a/management/server/networks/resources/manager.go +++ b/management/server/networks/resources/manager.go @@ -169,7 +169,10 @@ func (m *managerImpl) CreateResource(ctx context.Context, userID string, resourc } if affectedPeerIDs := m.resolveResourceAffectedPeers(ctx, resource.AccountID, affectedData); len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("CreateResource %s: updating %d affected peers: %v", resource.ID, len(affectedPeerIDs), affectedPeerIDs) go m.accountManager.UpdateAffectedPeers(ctx, resource.AccountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("CreateResource %s: no affected peers", resource.ID) } return resource, nil @@ -294,7 +297,10 @@ func (m *managerImpl) UpdateResource(ctx context.Context, userID string, resourc }() if affectedPeerIDs := m.resolveResourceAffectedPeers(ctx, resource.AccountID, affectedData); len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("UpdateResource %s: updating %d affected peers: %v", resource.ID, len(affectedPeerIDs), affectedPeerIDs) go m.accountManager.UpdateAffectedPeers(ctx, resource.AccountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("UpdateResource %s: no affected peers", resource.ID) } return resource, nil @@ -393,7 +399,10 @@ func (m *managerImpl) DeleteResource(ctx context.Context, accountID, userID, net } if affectedPeerIDs := m.resolveResourceAffectedPeers(ctx, accountID, affectedData); len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("DeleteResource %s: updating %d affected peers: %v", resourceID, len(affectedPeerIDs), affectedPeerIDs) go m.accountManager.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("DeleteResource %s: no affected peers", resourceID) } return nil @@ -491,6 +500,8 @@ func (m *managerImpl) resolveResourceAffectedPeers(ctx context.Context, accountI return nil } + log.WithContext(ctx).Tracef("resolveResourceAffectedPeers: resourceGroupIDs=%v, routerPeerGroups=%v, routerDirectPeers=%v, policies=%d", + data.resourceGroupIDs, data.routerPeerGroups, data.routerDirectPeers, len(data.policies)) groupSet := make(map[string]struct{}) var directPeerIDs []string @@ -559,6 +570,7 @@ func (m *managerImpl) resolveResourceAffectedPeers(ctx context.Context, accountI } } + log.WithContext(ctx).Tracef("resolveResourceAffectedPeers: result %d peers: %v", len(peerIDs), peerIDs) return peerIDs } diff --git a/management/server/networks/routers/manager.go b/management/server/networks/routers/manager.go index 6b33f558e..dc15fe974 100644 --- a/management/server/networks/routers/manager.go +++ b/management/server/networks/routers/manager.go @@ -128,7 +128,10 @@ func (m *managerImpl) CreateRouter(ctx context.Context, userID string, router *t m.accountManager.StoreEvent(ctx, userID, router.ID, router.AccountID, activity.NetworkRouterCreated, router.EventMeta(network)) if affectedPeerIDs := m.resolveRouterAffectedPeers(ctx, router.AccountID, affectedData); len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("CreateRouter %s: updating %d affected peers: %v", router.ID, len(affectedPeerIDs), affectedPeerIDs) go m.accountManager.UpdateAffectedPeers(ctx, router.AccountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("CreateRouter %s: no affected peers", router.ID) } return router, nil @@ -213,7 +216,10 @@ func (m *managerImpl) UpdateRouter(ctx context.Context, userID string, router *t m.accountManager.StoreEvent(ctx, userID, router.ID, router.AccountID, activity.NetworkRouterUpdated, router.EventMeta(network)) if affectedPeerIDs := m.resolveRouterAffectedPeers(ctx, router.AccountID, affectedData); len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("UpdateRouter %s: updating %d affected peers: %v", router.ID, len(affectedPeerIDs), affectedPeerIDs) go m.accountManager.UpdateAffectedPeers(ctx, router.AccountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("UpdateRouter %s: no affected peers", router.ID) } return router, nil @@ -261,7 +267,10 @@ func (m *managerImpl) DeleteRouter(ctx context.Context, accountID, userID, netwo event() if affectedPeerIDs := m.resolveRouterAffectedPeers(ctx, accountID, affectedData); len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("DeleteRouter %s: updating %d affected peers: %v", routerID, len(affectedPeerIDs), affectedPeerIDs) go m.accountManager.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("DeleteRouter %s: no affected peers", routerID) } return nil @@ -356,6 +365,8 @@ func (m *managerImpl) resolveRouterAffectedPeers(ctx context.Context, accountID return nil } + log.WithContext(ctx).Tracef("resolveRouterAffectedPeers: routerPeerGroups=%v, directPeerIDs=%v, resourceGroupIDs=%v, policies=%d", + data.routerPeerGroups, data.directPeerIDs, data.resourceGroupIDs, len(data.policies)) groupSet := make(map[string]struct{}) for _, gID := range data.routerPeerGroups { @@ -421,6 +432,7 @@ func (m *managerImpl) resolveRouterAffectedPeers(ctx context.Context, accountID } } + log.WithContext(ctx).Tracef("resolveRouterAffectedPeers: result %d peers: %v", len(peerIDs), peerIDs) return peerIDs } diff --git a/management/server/peer.go b/management/server/peer.go index 08d5b6deb..079a6389e 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -1308,7 +1308,7 @@ func (am *DefaultAccountManager) resolvePeerIDs(ctx context.Context, s store.Sto } if len(directPeerIDs) == 0 { - log.WithContext(ctx).Tracef("resolvePeerIDs: groups=%v -> %d peers", groupIDs, len(peerIDs)) + log.WithContext(ctx).Tracef("resolvePeerIDs: groups=%v -> %d peers: %v", groupIDs, len(peerIDs), peerIDs) return peerIDs } @@ -1323,7 +1323,7 @@ func (am *DefaultAccountManager) resolvePeerIDs(ctx context.Context, s store.Sto } } - log.WithContext(ctx).Tracef("resolvePeerIDs: groups=%v + directPeers=%v -> %d peers", groupIDs, directPeerIDs, len(peerIDs)) + log.WithContext(ctx).Tracef("resolvePeerIDs: groups=%v + directPeers=%v -> %d peers: %v", groupIDs, directPeerIDs, len(peerIDs), peerIDs) return peerIDs } diff --git a/management/server/policy.go b/management/server/policy.go index 0404a9208..6c053f5fb 100644 --- a/management/server/policy.go +++ b/management/server/policy.go @@ -5,7 +5,7 @@ import ( _ "embed" "github.com/rs/xid" - "github.com/sirupsen/logrus" + log "github.com/sirupsen/logrus" "github.com/netbirdio/netbird/management/server/permissions/modules" "github.com/netbirdio/netbird/management/server/permissions/operations" @@ -58,7 +58,7 @@ func (am *DefaultAccountManager) SavePolicy(ctx context.Context, accountID, user if isUpdate { if policy.Equal(existingPolicy) { - logrus.WithContext(ctx).Tracef("policy update skipped because equal to stored one - policy id %s", policy.ID) + log.WithContext(ctx).Tracef("policy update skipped because equal to stored one - policy id %s", policy.ID) unchanged = true return nil } @@ -74,7 +74,7 @@ func (am *DefaultAccountManager) SavePolicy(ctx context.Context, accountID, user } } - groupIDs, directPeerIDs := collectPolicyAffectedGroupsAndPeers(policy, existingPolicy) + groupIDs, directPeerIDs := collectPolicyAffectedGroupsAndPeers(ctx, policy, existingPolicy) affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) return transaction.IncrementNetworkSerial(ctx, accountID) @@ -90,7 +90,10 @@ func (am *DefaultAccountManager) SavePolicy(ctx context.Context, accountID, user am.StoreEvent(ctx, userID, policy.ID, accountID, action, policy.EventMeta()) if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Tracef("SavePolicy %s: updating %d affected peers: %v", policy.ID, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("SavePolicy %s: no affected peers", policy.ID) } return policy, nil @@ -115,7 +118,7 @@ func (am *DefaultAccountManager) DeletePolicy(ctx context.Context, accountID, po return err } - groupIDs, directPeerIDs := collectPolicyAffectedGroupsAndPeers(policy) + groupIDs, directPeerIDs := collectPolicyAffectedGroupsAndPeers(ctx, policy) affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) if err = transaction.DeletePolicy(ctx, accountID, policyID); err != nil { @@ -131,7 +134,10 @@ func (am *DefaultAccountManager) DeletePolicy(ctx context.Context, accountID, po am.StoreEvent(ctx, userID, policyID, accountID, activity.PolicyRemoved, policy.EventMeta()) if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("DeletePolicy %s: updating %d affected peers: %v", policyID, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("DeletePolicy %s: no affected peers", policyID) } return nil @@ -151,21 +157,26 @@ func (am *DefaultAccountManager) ListPolicies(ctx context.Context, accountID, us } // collectPolicyAffectedGroupsAndPeers returns group IDs and direct peer IDs from the given policies. -func collectPolicyAffectedGroupsAndPeers(policies ...*types.Policy) (groupIDs []string, directPeerIDs []string) { +func collectPolicyAffectedGroupsAndPeers(ctx context.Context, policies ...*types.Policy) (groupIDs []string, directPeerIDs []string) { for _, policy := range policies { if policy == nil { continue } - groupIDs = append(groupIDs, policy.RuleGroups()...) + ruleGroups := policy.RuleGroups() + log.WithContext(ctx).Tracef("collectPolicyAffectedGroupsAndPeers: policy %s (%s) ruleGroups=%v", policy.ID, policy.Name, ruleGroups) + groupIDs = append(groupIDs, ruleGroups...) for _, rule := range policy.Rules { if rule.SourceResource.Type == types.ResourceTypePeer && rule.SourceResource.ID != "" { + log.WithContext(ctx).Tracef("collectPolicyAffectedGroupsAndPeers: policy %s rule %s direct source peer %s", policy.ID, rule.ID, rule.SourceResource.ID) directPeerIDs = append(directPeerIDs, rule.SourceResource.ID) } if rule.DestinationResource.Type == types.ResourceTypePeer && rule.DestinationResource.ID != "" { + log.WithContext(ctx).Tracef("collectPolicyAffectedGroupsAndPeers: policy %s rule %s direct destination peer %s", policy.ID, rule.ID, rule.DestinationResource.ID) directPeerIDs = append(directPeerIDs, rule.DestinationResource.ID) } } } + log.WithContext(ctx).Tracef("collectPolicyAffectedGroupsAndPeers: result groupIDs=%v, directPeerIDs=%v", groupIDs, directPeerIDs) return } diff --git a/management/server/posture_checks.go b/management/server/posture_checks.go index f22bc8d14..8377ce58b 100644 --- a/management/server/posture_checks.go +++ b/management/server/posture_checks.go @@ -75,7 +75,10 @@ func (am *DefaultAccountManager) SavePostureChecks(ctx context.Context, accountI am.StoreEvent(ctx, userID, postureChecks.ID, accountID, action, postureChecks.EventMeta()) if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("SavePostureChecks %s: updating %d affected peers: %v", postureChecks.ID, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("SavePostureChecks %s: no affected peers", postureChecks.ID) } return postureChecks, nil @@ -141,12 +144,14 @@ func collectPostureCheckAffectedGroupsAndPeers(ctx context.Context, transaction for _, policy := range policies { if slices.Contains(policy.SourcePostureChecks, postureCheckID) { - gIDs, pIDs := collectPolicyAffectedGroupsAndPeers(policy) + log.WithContext(ctx).Tracef("collectPostureCheckAffectedGroupsAndPeers: posture check %s referenced by policy %s (%s)", postureCheckID, policy.ID, policy.Name) + gIDs, pIDs := collectPolicyAffectedGroupsAndPeers(ctx, policy) groupIDs = append(groupIDs, gIDs...) directPeerIDs = append(directPeerIDs, pIDs...) } } + log.WithContext(ctx).Tracef("collectPostureCheckAffectedGroupsAndPeers: postureCheck=%s -> groupIDs=%v, directPeerIDs=%v", postureCheckID, groupIDs, directPeerIDs) return groupIDs, directPeerIDs } diff --git a/management/server/route.go b/management/server/route.go index 30297f851..3a39518f2 100644 --- a/management/server/route.go +++ b/management/server/route.go @@ -8,6 +8,7 @@ import ( "unicode/utf8" "github.com/rs/xid" + log "github.com/sirupsen/logrus" "github.com/netbirdio/netbird/management/server/activity" "github.com/netbirdio/netbird/management/server/permissions/modules" @@ -177,7 +178,7 @@ func (am *DefaultAccountManager) CreateRoute(ctx context.Context, accountID stri return err } - groupIDs, directPeerIDs := collectRouteAffectedGroupsAndPeers(newRoute) + groupIDs, directPeerIDs := collectRouteAffectedGroupsAndPeers(ctx, newRoute) affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) return transaction.IncrementNetworkSerial(ctx, accountID) @@ -189,7 +190,10 @@ func (am *DefaultAccountManager) CreateRoute(ctx context.Context, accountID stri am.StoreEvent(ctx, userID, string(newRoute.ID), accountID, activity.RouteCreated, newRoute.EventMeta()) if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("CreateRoute %s: updating %d affected peers: %v", newRoute.ID, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("CreateRoute %s: no affected peers", newRoute.ID) } return newRoute, nil @@ -224,7 +228,7 @@ func (am *DefaultAccountManager) SaveRoute(ctx context.Context, accountID, userI return err } - groupIDs, directPeerIDs := collectRouteAffectedGroupsAndPeers(routeToSave, oldRoute) + groupIDs, directPeerIDs := collectRouteAffectedGroupsAndPeers(ctx, routeToSave, oldRoute) affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) return transaction.IncrementNetworkSerial(ctx, accountID) @@ -236,7 +240,10 @@ func (am *DefaultAccountManager) SaveRoute(ctx context.Context, accountID, userI am.StoreEvent(ctx, userID, string(routeToSave.ID), accountID, activity.RouteUpdated, routeToSave.EventMeta()) if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("SaveRoute %s: updating %d affected peers: %v", routeToSave.ID, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("SaveRoute %s: no affected peers", routeToSave.ID) } return nil @@ -261,7 +268,7 @@ func (am *DefaultAccountManager) DeleteRoute(ctx context.Context, accountID stri return err } - groupIDs, directPeerIDs := collectRouteAffectedGroupsAndPeers(rt) + groupIDs, directPeerIDs := collectRouteAffectedGroupsAndPeers(ctx, rt) affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) if err = transaction.DeleteRoute(ctx, accountID, string(routeID)); err != nil { @@ -277,7 +284,10 @@ func (am *DefaultAccountManager) DeleteRoute(ctx context.Context, accountID stri am.StoreEvent(ctx, userID, string(rt.ID), accountID, activity.RouteRemoved, rt.EventMeta()) if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("DeleteRoute %s: updating %d affected peers: %v", routeID, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("DeleteRoute %s: no affected peers", routeID) } return nil @@ -367,11 +377,13 @@ func getPlaceholderIP() netip.Prefix { } // collectRouteAffectedGroupsAndPeers returns group IDs and direct peer IDs from the given routes. -func collectRouteAffectedGroupsAndPeers(routes ...*route.Route) (groupIDs []string, directPeerIDs []string) { +func collectRouteAffectedGroupsAndPeers(ctx context.Context, routes ...*route.Route) (groupIDs []string, directPeerIDs []string) { for _, r := range routes { if r == nil { continue } + log.WithContext(ctx).Tracef("collectRouteAffectedGroupsAndPeers: route %s groups=%v peerGroups=%v accessControlGroups=%v peer=%q", + r.ID, r.Groups, r.PeerGroups, r.AccessControlGroups, r.Peer) groupIDs = append(groupIDs, r.Groups...) groupIDs = append(groupIDs, r.PeerGroups...) groupIDs = append(groupIDs, r.AccessControlGroups...) @@ -379,6 +391,7 @@ func collectRouteAffectedGroupsAndPeers(routes ...*route.Route) (groupIDs []stri directPeerIDs = append(directPeerIDs, r.Peer) } } + log.WithContext(ctx).Tracef("collectRouteAffectedGroupsAndPeers: result groupIDs=%v, directPeerIDs=%v", groupIDs, directPeerIDs) return } From 40e6ec16c6153bf4d50ee308265f3d5e76f53742 Mon Sep 17 00:00:00 2001 From: pascal Date: Thu, 7 May 2026 17:36:09 +0200 Subject: [PATCH 11/28] log --- management/server/group.go | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/management/server/group.go b/management/server/group.go index 5d15eed39..cce7bea90 100644 --- a/management/server/group.go +++ b/management/server/group.go @@ -894,8 +894,9 @@ func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Sto if !policyReferencesGroups(policy, changedSet) { continue } - log.WithContext(ctx).Tracef("policy %s (%s) references changed groups, adding rule groups", policy.ID, policy.Name) - for _, gID := range policy.RuleGroups() { + ruleGroups := policy.RuleGroups() + log.WithContext(ctx).Tracef("policy %s (%s) references changed groups, adding rule groups %v", policy.ID, policy.Name, ruleGroups) + for _, gID := range ruleGroups { groupSet[gID] = struct{}{} } for _, rule := range policy.Rules { From 57529c7f185052976bc383e4601877940c40bc41 Mon Sep 17 00:00:00 2001 From: pascal Date: Thu, 7 May 2026 17:50:02 +0200 Subject: [PATCH 12/28] linter --- .../network_map/controller/controller.go | 39 ------------------- management/server/affected_peers_test.go | 8 ++-- .../server/networks/resources/manager.go | 2 +- management/server/networks/routers/manager.go | 2 +- 4 files changed, 6 insertions(+), 45 deletions(-) diff --git a/management/internals/controllers/network_map/controller/controller.go b/management/internals/controllers/network_map/controller/controller.go index b17036c6b..d5be1fd65 100644 --- a/management/internals/controllers/network_map/controller/controller.go +++ b/management/internals/controllers/network_map/controller/controller.go @@ -44,7 +44,6 @@ type Controller struct { EphemeralPeersManager ephemeral.Manager accountUpdateLocks sync.Map - sendAccountUpdateLocks sync.Map affectedPeerUpdateLocks sync.Map updateAccountPeersBufferInterval atomic.Int64 // dnsDomain is used for peer resolution. This is appended to the peer's name @@ -229,44 +228,6 @@ func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID strin return nil } -func (c *Controller) bufferSendUpdateAccountPeers(ctx context.Context, accountID string, reason types.UpdateReason) error { - log.WithContext(ctx).Tracef("buffer sending update peers for account %s from %s", accountID, util.GetCallerName()) - - if c.accountManagerMetrics != nil { - c.accountManagerMetrics.CountUpdateAccountPeersTriggered(string(reason.Resource), string(reason.Operation)) - } - - bufUpd, _ := c.sendAccountUpdateLocks.LoadOrStore(accountID, &bufferUpdate{}) - b := bufUpd.(*bufferUpdate) - - if !b.mu.TryLock() { - b.update.Store(true) - return nil - } - - if b.next != nil { - b.next.Stop() - } - - go func() { - defer b.mu.Unlock() - _ = c.sendUpdateAccountPeers(ctx, accountID) - if !b.update.Load() { - return - } - b.update.Store(false) - if b.next == nil { - b.next = time.AfterFunc(time.Duration(c.updateAccountPeersBufferInterval.Load()), func() { - _ = c.sendUpdateAccountPeers(ctx, accountID) - }) - return - } - b.next.Reset(time.Duration(c.updateAccountPeersBufferInterval.Load())) - }() - - return nil -} - // UpdatePeers updates all peers that belong to an account. // Should be called when changes have to be synced to peers. func (c *Controller) UpdateAccountPeers(ctx context.Context, accountID string, reason types.UpdateReason) error { diff --git a/management/server/affected_peers_test.go b/management/server/affected_peers_test.go index 13bb990c8..3208a7ec0 100644 --- a/management/server/affected_peers_test.go +++ b/management/server/affected_peers_test.go @@ -481,12 +481,12 @@ func TestCollectPolicyAffectedGroups_MultiplePolicies(t *testing.T) { {Sources: []string{"g1"}, Destinations: []string{"g2"}}, }, } - new := &types.Policy{ + updated := &types.Policy{ Rules: []*types.PolicyRule{ {Sources: []string{"g3"}, Destinations: []string{"g4"}}, }, } - groups, _ := collectPolicyAffectedGroupsAndPeers(context.Background(), new, old) + groups, _ := collectPolicyAffectedGroupsAndPeers(context.Background(), updated, old) assert.ElementsMatch(t, []string{"g1", "g2", "g3", "g4"}, groups) } @@ -548,11 +548,11 @@ func TestCollectRouteAffectedGroups_MultipleRoutes(t *testing.T) { Groups: []string{"g1"}, Peer: "p1", } - new := &route.Route{ + updated := &route.Route{ Groups: []string{"g2"}, PeerGroups: []string{"g3"}, } - groups, directPeers := collectRouteAffectedGroupsAndPeers(context.Background(), new, old) + groups, directPeers := collectRouteAffectedGroupsAndPeers(context.Background(), updated, old) assert.ElementsMatch(t, []string{"g1", "g2", "g3"}, groups) assert.ElementsMatch(t, []string{"p1"}, directPeers) } diff --git a/management/server/networks/resources/manager.go b/management/server/networks/resources/manager.go index 23f4f98a4..03dbe542b 100644 --- a/management/server/networks/resources/manager.go +++ b/management/server/networks/resources/manager.go @@ -461,7 +461,7 @@ type resourceAffectedPeersData struct { // loadResourceAffectedPeersData loads the data needed to determine affected peers within a transaction. func loadResourceAffectedPeersData(ctx context.Context, transaction store.Store, accountID, networkID string, resourceGroupIDs []string) (*resourceAffectedPeersData, error) { if len(resourceGroupIDs) == 0 { - return nil, nil + return &resourceAffectedPeersData{}, nil } policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) diff --git a/management/server/networks/routers/manager.go b/management/server/networks/routers/manager.go index dc15fe974..0da184ca3 100644 --- a/management/server/networks/routers/manager.go +++ b/management/server/networks/routers/manager.go @@ -321,7 +321,7 @@ func loadRouterAffectedPeersData(ctx context.Context, transaction store.Store, a } if len(routerPeerGroups) == 0 && len(directPeerIDs) == 0 { - return nil, nil + return &routerAffectedPeersData{}, nil } resources, err := transaction.GetNetworkResourcesByNetID(ctx, store.LockingStrengthNone, accountID, networkID) From 70e84d5228300600f8f9c2d7520437e704bad91b Mon Sep 17 00:00:00 2001 From: pascal Date: Thu, 7 May 2026 18:07:47 +0200 Subject: [PATCH 13/28] add own peer on peer update --- management/server/account_test.go | 10 ++++++++ management/server/peer.go | 1 + management/server/peer_test.go | 39 +++++++++++++++++-------------- 3 files changed, 32 insertions(+), 18 deletions(-) diff --git a/management/server/account_test.go b/management/server/account_test.go index 6bb875f99..7c86aebf8 100644 --- a/management/server/account_test.go +++ b/management/server/account_test.go @@ -3203,6 +3203,16 @@ func setupNetworkMapTest(t *testing.T) (*DefaultAccountManager, *update_channel. // when the channel delivers. const peerUpdateTimeout = 5 * time.Second +func drainPeerUpdates(ch <-chan *network_map.UpdateMessage) { + for { + select { + case <-ch: + case <-time.After(200 * time.Millisecond): + return + } + } +} + func peerShouldNotReceiveUpdate(t *testing.T, updateMessage <-chan *network_map.UpdateMessage) { t.Helper() select { diff --git a/management/server/peer.go b/management/server/peer.go index 079a6389e..5ea0197bc 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -299,6 +299,7 @@ func (am *DefaultAccountManager) UpdatePeer(ctx context.Context, accountID, user changedPeerIDs := []string{peer.ID} affectedPeerIDs := am.resolveAffectedPeersForPeerChanges(ctx, am.Store, accountID, changedPeerIDs) + affectedPeerIDs = append(affectedPeerIDs, peer.ID) err = am.networkMapController.OnPeersUpdated(ctx, accountID, changedPeerIDs, affectedPeerIDs) if err != nil { return nil, fmt.Errorf("notify network map controller of peer update: %w", err) diff --git a/management/server/peer_test.go b/management/server/peer_test.go index 07acf865f..36af9af81 100644 --- a/management/server/peer_test.go +++ b/management/server/peer_test.go @@ -1855,7 +1855,7 @@ func TestPeerAccountPeersUpdate(t *testing.T) { t.Run("adding peer to unlinked group", func(t *testing.T) { done := make(chan struct{}) go func() { - peerShouldReceiveUpdate(t, updMsg) // + peerShouldNotReceiveUpdate(t, updMsg) close(done) }() @@ -1880,7 +1880,7 @@ func TestPeerAccountPeersUpdate(t *testing.T) { t.Run("deleting peer with unlinked group", func(t *testing.T) { done := make(chan struct{}) go func() { - peerShouldReceiveUpdate(t, updMsg) + peerShouldNotReceiveUpdate(t, updMsg) close(done) }() @@ -2018,7 +2018,10 @@ func TestPeerAccountPeersUpdate(t *testing.T) { } }) - // Adding peer to group linked with route should update account peers and send peer update + // drain any buffered updates from previous subtests + drainPeerUpdates(updMsg) + + // Adding peer to group linked with route should update peers in that group, not unrelated peers t.Run("adding peer to group linked with route", func(t *testing.T) { route := nbroute.Route{ ID: "testingRoute1", @@ -2042,7 +2045,7 @@ func TestPeerAccountPeersUpdate(t *testing.T) { done := make(chan struct{}) go func() { - peerShouldReceiveUpdate(t, updMsg) + peerShouldNotReceiveUpdate(t, updMsg) close(done) }() @@ -2059,16 +2062,16 @@ func TestPeerAccountPeersUpdate(t *testing.T) { select { case <-done: - case <-time.After(peerUpdateTimeout): - t.Error("timeout waiting for peerShouldReceiveUpdate") + case <-time.After(time.Second): + t.Error("timeout waiting for peerShouldNotReceiveUpdate") } }) - // Deleting peer with linked group to route should update account peers and send peer update + // Deleting peer with linked group to route should update peers in that group, not unrelated peers t.Run("deleting peer with linked group to route", func(t *testing.T) { done := make(chan struct{}) go func() { - peerShouldReceiveUpdate(t, updMsg) + peerShouldNotReceiveUpdate(t, updMsg) close(done) }() @@ -2077,12 +2080,12 @@ func TestPeerAccountPeersUpdate(t *testing.T) { select { case <-done: - case <-time.After(peerUpdateTimeout): - t.Error("timeout waiting for peerShouldReceiveUpdate") + case <-time.After(time.Second): + t.Error("timeout waiting for peerShouldNotReceiveUpdate") } }) - // Adding peer to group linked with name server group should update account peers and send peer update + // Adding peer to group linked with name server group should update peers in that group, not unrelated peers t.Run("adding peer to group linked with name server group", func(t *testing.T) { _, err = manager.CreateNameServerGroup( context.Background(), account.Id, "nsGroup", "nsGroup", []nbdns.NameServer{{ @@ -2097,7 +2100,7 @@ func TestPeerAccountPeersUpdate(t *testing.T) { done := make(chan struct{}) go func() { - peerShouldReceiveUpdate(t, updMsg) + peerShouldNotReceiveUpdate(t, updMsg) close(done) }() @@ -2114,16 +2117,16 @@ func TestPeerAccountPeersUpdate(t *testing.T) { select { case <-done: - case <-time.After(peerUpdateTimeout): - t.Error("timeout waiting for peerShouldReceiveUpdate") + case <-time.After(time.Second): + t.Error("timeout waiting for peerShouldNotReceiveUpdate") } }) - // Deleting peer with linked group to name server group should update account peers and send peer update + // Deleting peer with linked group to name server group should update peers in that group, not unrelated peers t.Run("deleting peer with linked group to route", func(t *testing.T) { done := make(chan struct{}) go func() { - peerShouldReceiveUpdate(t, updMsg) + peerShouldNotReceiveUpdate(t, updMsg) close(done) }() @@ -2132,8 +2135,8 @@ func TestPeerAccountPeersUpdate(t *testing.T) { select { case <-done: - case <-time.After(peerUpdateTimeout): - t.Error("timeout waiting for peerShouldReceiveUpdate") + case <-time.After(time.Second): + t.Error("timeout waiting for peerShouldNotReceiveUpdate") } }) } From fed4f1b0241be41f565e791f911a475c1b6b33fc Mon Sep 17 00:00:00 2001 From: pascal Date: Fri, 8 May 2026 14:33:31 +0200 Subject: [PATCH 14/28] drain channel between tests --- management/server/user_test.go | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/management/server/user_test.go b/management/server/user_test.go index 68fc58eef..5424e2fef 100644 --- a/management/server/user_test.go +++ b/management/server/user_test.go @@ -1531,11 +1531,14 @@ func TestUserAccountPeersUpdate(t *testing.T) { } }) + // drain any buffered updates from previous subtests + drainPeerUpdates(updMsg) + // deleting user with no linked peers should not update account peers and not send peer update t.Run("deleting user with no linked peers", func(t *testing.T) { done := make(chan struct{}) go func() { - peerShouldReceiveUpdate(t, updMsg) + peerShouldNotReceiveUpdate(t, updMsg) close(done) }() From 85851bc4779193bcf50d54c6a7414a49b677844a Mon Sep 17 00:00:00 2001 From: pascal Date: Fri, 8 May 2026 16:43:27 +0200 Subject: [PATCH 15/28] extract submethods --- .../network_map/controller/controller.go | 46 ++- management/server/group.go | 390 ------------------ management/server/networks/manager.go | 90 ++-- .../server/networks/resources/manager.go | 198 +++++---- management/server/networks/routers/manager.go | 176 ++++---- 5 files changed, 288 insertions(+), 612 deletions(-) diff --git a/management/internals/controllers/network_map/controller/controller.go b/management/internals/controllers/network_map/controller/controller.go index d5be1fd65..455db9a74 100644 --- a/management/internals/controllers/network_map/controller/controller.go +++ b/management/internals/controllers/network_map/controller/controller.go @@ -248,19 +248,7 @@ func (c *Controller) UpdateAffectedPeers(ctx context.Context, accountID string, func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID string, peerIDs []string) error { log.WithContext(ctx).Tracef("sendUpdateForAffectedPeers: account %s, %d affected peers: %v (caller: %s)", accountID, len(peerIDs), peerIDs, util.GetCallerName()) - affected := make(map[string]struct{}, len(peerIDs)) - for _, id := range peerIDs { - affected[id] = struct{}{} - } - - hasConnected := false - for _, id := range peerIDs { - if c.peersUpdateManager.HasChannel(id) { - hasConnected = true - break - } - } - if !hasConnected { + if !c.hasConnectedPeers(peerIDs) { log.WithContext(ctx).Tracef("sendUpdateForAffectedPeers: no connected peers among %v, skipping", peerIDs) return nil } @@ -272,13 +260,7 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s globalStart := time.Now() - var peersToUpdate []*nbpeer.Peer - for _, peer := range account.Peers { - if _, ok := affected[peer.ID]; ok && c.peersUpdateManager.HasChannel(peer.ID) { - peersToUpdate = append(peersToUpdate, peer) - } - } - + peersToUpdate := c.filterConnectedAffectedPeers(account, peerIDs) if len(peersToUpdate) == 0 { log.WithContext(ctx).Tracef("sendUpdateForAffectedPeers: no peers to update (affected peers not found in account or no channels)") return nil @@ -368,6 +350,30 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s return nil } +func (c *Controller) hasConnectedPeers(peerIDs []string) bool { + for _, id := range peerIDs { + if c.peersUpdateManager.HasChannel(id) { + return true + } + } + return false +} + +func (c *Controller) filterConnectedAffectedPeers(account *types.Account, peerIDs []string) []*nbpeer.Peer { + affected := make(map[string]struct{}, len(peerIDs)) + for _, id := range peerIDs { + affected[id] = struct{}{} + } + + var result []*nbpeer.Peer + for _, peer := range account.Peers { + if _, ok := affected[peer.ID]; ok && c.peersUpdateManager.HasChannel(peer.ID) { + result = append(result, peer) + } + } + return result +} + func (c *Controller) UpdateAccountPeer(ctx context.Context, accountId string, peerId string) error { if !c.peersUpdateManager.HasChannel(peerId) { return fmt.Errorf("peer %s doesn't have a channel, skipping network map update", peerId) diff --git a/management/server/group.go b/management/server/group.go index cce7bea90..125c11374 100644 --- a/management/server/group.go +++ b/management/server/group.go @@ -9,15 +9,12 @@ import ( "github.com/rs/xid" log "github.com/sirupsen/logrus" - nbdns "github.com/netbirdio/netbird/dns" "github.com/netbirdio/netbird/management/server/activity" - routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" "github.com/netbirdio/netbird/management/server/permissions/modules" "github.com/netbirdio/netbird/management/server/permissions/operations" "github.com/netbirdio/netbird/management/server/store" "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/management/server/util" - "github.com/netbirdio/netbird/route" "github.com/netbirdio/netbird/shared/management/status" ) @@ -656,390 +653,3 @@ func validateNewGroup(ctx context.Context, transaction store.Store, accountID st return nil } - -func validateDeleteGroup(ctx context.Context, transaction store.Store, group *types.Group, userID string, flowGroups []string) error { - // disable a deleting integration group if the initiator is not an admin service user - if group.Issued == types.GroupIssuedIntegration { - executingUser, err := transaction.GetUserByUserID(ctx, store.LockingStrengthNone, userID) - if err != nil { - return status.Errorf(status.Internal, "failed to get user") - } - if executingUser.Role != types.UserRoleAdmin || !executingUser.IsServiceUser { - return status.Errorf(status.PermissionDenied, "only service users with admin power can delete integration group") - } - } - - if group.IsGroupAll() { - return status.Errorf(status.InvalidArgument, "deleting group ALL is not allowed") - } - - if len(group.Resources) > 0 { - return &GroupLinkError{"network resource", group.Resources[0].ID} - } - - if slices.Contains(flowGroups, group.ID) { - return &GroupLinkError{"settings", "traffic event logging"} - } - - if isLinked, linkedRoute := isGroupLinkedToRoute(ctx, transaction, group.AccountID, group.ID); isLinked { - return &GroupLinkError{"route", string(linkedRoute.NetID)} - } - - if isLinked, linkedDns := isGroupLinkedToDns(ctx, transaction, group.AccountID, group.ID); isLinked { - return &GroupLinkError{"name server groups", linkedDns.Name} - } - - if isLinked, linkedPolicy := isGroupLinkedToPolicy(ctx, transaction, group.AccountID, group.ID); isLinked { - return &GroupLinkError{"policy", linkedPolicy.Name} - } - - if isLinked, linkedSetupKey := isGroupLinkedToSetupKey(ctx, transaction, group.AccountID, group.ID); isLinked { - return &GroupLinkError{"setup key", linkedSetupKey.Name} - } - - if isLinked, linkedUser := isGroupLinkedToUser(ctx, transaction, group.AccountID, group.ID); isLinked { - return &GroupLinkError{"user", linkedUser.Id} - } - - if isLinked, linkedRouter := isGroupLinkedToNetworkRouter(ctx, transaction, group.AccountID, group.ID); isLinked { - return &GroupLinkError{"network router", linkedRouter.ID} - } - - return checkGroupLinkedToSettings(ctx, transaction, group) -} - -// checkGroupLinkedToSettings verifies if a group is linked to any settings in the account. -func checkGroupLinkedToSettings(ctx context.Context, transaction store.Store, group *types.Group) error { - dnsSettings, err := transaction.GetAccountDNSSettings(ctx, store.LockingStrengthNone, group.AccountID) - if err != nil { - return status.Errorf(status.Internal, "failed to get DNS settings") - } - - if slices.Contains(dnsSettings.DisabledManagementGroups, group.ID) { - return &GroupLinkError{"disabled DNS management groups", group.Name} - } - - settings, err := transaction.GetAccountSettings(ctx, store.LockingStrengthNone, group.AccountID) - if err != nil { - return status.Errorf(status.Internal, "failed to get account settings") - } - - if settings.Extra != nil && slices.Contains(settings.Extra.IntegratedValidatorGroups, group.ID) { - return &GroupLinkError{"integrated validator", group.Name} - } - - return nil -} - -// isGroupLinkedToRoute checks if a group is linked to any route in the account. -func isGroupLinkedToRoute(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *route.Route) { - routes, err := transaction.GetAccountRoutes(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("error retrieving routes while checking group linkage: %v", err) - return false, nil - } - - for _, r := range routes { - isLinked := slices.Contains(r.Groups, groupID) || - slices.Contains(r.PeerGroups, groupID) || - slices.Contains(r.AccessControlGroups, groupID) - if isLinked { - return true, r - } - } - - return false, nil -} - -// isGroupLinkedToPolicy checks if a group is linked to any policy in the account. -func isGroupLinkedToPolicy(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *types.Policy) { - policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("error retrieving policies while checking group linkage: %v", err) - return false, nil - } - - for _, policy := range policies { - for _, rule := range policy.Rules { - if slices.Contains(rule.Sources, groupID) || slices.Contains(rule.Destinations, groupID) { - return true, policy - } - } - } - return false, nil -} - -// isGroupLinkedToDns checks if a group is linked to any nameserver group in the account. -func isGroupLinkedToDns(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *nbdns.NameServerGroup) { - nameServerGroups, err := transaction.GetAccountNameServerGroups(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("error retrieving name server groups while checking group linkage: %v", err) - return false, nil - } - - for _, dns := range nameServerGroups { - for _, g := range dns.Groups { - if g == groupID { - return true, dns - } - } - } - - return false, nil -} - -// isGroupLinkedToSetupKey checks if a group is linked to any setup key in the account. -func isGroupLinkedToSetupKey(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *types.SetupKey) { - setupKeys, err := transaction.GetAccountSetupKeys(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("error retrieving setup keys while checking group linkage: %v", err) - return false, nil - } - - for _, setupKey := range setupKeys { - if slices.Contains(setupKey.AutoGroups, groupID) { - return true, setupKey - } - } - return false, nil -} - -// isGroupLinkedToUser checks if a group is linked to any user in the account. -func isGroupLinkedToUser(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *types.User) { - users, err := transaction.GetAccountUsers(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("error retrieving users while checking group linkage: %v", err) - return false, nil - } - - for _, user := range users { - if slices.Contains(user.AutoGroups, groupID) { - return true, user - } - } - return false, nil -} - -// isGroupLinkedToNetworkRouter checks if a group is linked to any network router in the account. -func isGroupLinkedToNetworkRouter(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *routerTypes.NetworkRouter) { - routers, err := transaction.GetNetworkRoutersByAccountID(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("error retrieving network routers while checking group linkage: %v", err) - return false, nil - } - - for _, router := range routers { - if slices.Contains(router.PeerGroups, groupID) { - return true, router - } - } - return false, nil -} - -// areGroupChangesAffectPeers checks if any changes to the specified groups will affect peers. -func areGroupChangesAffectPeers(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) (bool, error) { - if len(groupIDs) == 0 { - return false, nil - } - - dnsSettings, err := transaction.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID) - if err != nil { - return false, err - } - - for _, groupID := range groupIDs { - if slices.Contains(dnsSettings.DisabledManagementGroups, groupID) { - return true, nil - } - if linked, _ := isGroupLinkedToDns(ctx, transaction, accountID, groupID); linked { - return true, nil - } - if linked, _ := isGroupLinkedToPolicy(ctx, transaction, accountID, groupID); linked { - return true, nil - } - if linked, _ := isGroupLinkedToRoute(ctx, transaction, accountID, groupID); linked { - return true, nil - } - if linked, _ := isGroupLinkedToNetworkRouter(ctx, transaction, accountID, groupID); linked { - return true, nil - } - } - - return false, nil -} - -// collectGroupChangeAffectedGroups walks policies, routes, nameservers, DNS settings, -// and network routers to collect all group IDs and direct peer IDs affected by the changed groups. -func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedGroupIDs []string) (allGroupIDs []string, directPeerIDs []string) { - if len(changedGroupIDs) == 0 { - return nil, nil - } - - changedSet := make(map[string]struct{}, len(changedGroupIDs)) - for _, id := range changedGroupIDs { - changedSet[id] = struct{}{} - } - - log.WithContext(ctx).Tracef("collecting affected groups for changed groups %v", changedGroupIDs) - - groupSet := make(map[string]struct{}) - - peerSet := make(map[string]struct{}) - - policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get policies for group change resolution: %v", err) - } else { - for _, policy := range policies { - if !policyReferencesGroups(policy, changedSet) { - continue - } - ruleGroups := policy.RuleGroups() - log.WithContext(ctx).Tracef("policy %s (%s) references changed groups, adding rule groups %v", policy.ID, policy.Name, ruleGroups) - for _, gID := range ruleGroups { - groupSet[gID] = struct{}{} - } - for _, rule := range policy.Rules { - if rule.SourceResource.Type == types.ResourceTypePeer && rule.SourceResource.ID != "" { - log.WithContext(ctx).Tracef("policy %s rule %s has direct source peer %s", policy.ID, rule.ID, rule.SourceResource.ID) - peerSet[rule.SourceResource.ID] = struct{}{} - } - if rule.DestinationResource.Type == types.ResourceTypePeer && rule.DestinationResource.ID != "" { - log.WithContext(ctx).Tracef("policy %s rule %s has direct destination peer %s", policy.ID, rule.ID, rule.DestinationResource.ID) - peerSet[rule.DestinationResource.ID] = struct{}{} - } - } - } - } - - routes, err := transaction.GetAccountRoutes(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get routes for group change resolution: %v", err) - } else { - for _, r := range routes { - if !routeReferencesGroups(r, changedSet) { - continue - } - log.WithContext(ctx).Tracef("route %s (%s) references changed groups", r.ID, r.Description) - for _, gID := range r.Groups { - groupSet[gID] = struct{}{} - } - for _, gID := range r.PeerGroups { - groupSet[gID] = struct{}{} - } - for _, gID := range r.AccessControlGroups { - groupSet[gID] = struct{}{} - } - if r.Peer != "" { - log.WithContext(ctx).Tracef("route %s has direct peer %s", r.ID, r.Peer) - peerSet[r.Peer] = struct{}{} - } - } - } - - nsGroups, err := transaction.GetAccountNameServerGroups(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get nameserver groups for group change resolution: %v", err) - } else { - for _, ns := range nsGroups { - for _, gID := range ns.Groups { - if _, ok := changedSet[gID]; ok { - log.WithContext(ctx).Tracef("nameserver group %s (%s) references changed group %s", ns.ID, ns.Name, gID) - for _, g := range ns.Groups { - groupSet[g] = struct{}{} - } - break - } - } - } - } - - dnsSettings, err := transaction.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get DNS settings for group change resolution: %v", err) - } else { - for _, gID := range dnsSettings.DisabledManagementGroups { - if _, ok := changedSet[gID]; ok { - log.WithContext(ctx).Tracef("DNS disabled management group %s matches changed group", gID) - groupSet[gID] = struct{}{} - } - } - } - - routers, err := transaction.GetNetworkRoutersByAccountID(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get network routers for group change resolution: %v", err) - } else { - for _, router := range routers { - if !routerReferencesGroups(router, changedSet) { - continue - } - log.WithContext(ctx).Tracef("network router %s references changed groups", router.ID) - for _, gID := range router.PeerGroups { - groupSet[gID] = struct{}{} - } - if router.Peer != "" { - log.WithContext(ctx).Tracef("network router %s has direct peer %s", router.ID, router.Peer) - peerSet[router.Peer] = struct{}{} - } - } - } - - allGroupIDs = make([]string, 0, len(groupSet)) - for gID := range groupSet { - allGroupIDs = append(allGroupIDs, gID) - } - - directPeerIDs = make([]string, 0, len(peerSet)) - for pID := range peerSet { - directPeerIDs = append(directPeerIDs, pID) - } - - log.WithContext(ctx).Tracef("affected groups resolution: changed=%v -> affectedGroups=%v, directPeers=%v", changedGroupIDs, allGroupIDs, directPeerIDs) - - return allGroupIDs, directPeerIDs -} - -func policyReferencesGroups(policy *types.Policy, groupSet map[string]struct{}) bool { - for _, rule := range policy.Rules { - for _, gID := range rule.Sources { - if _, ok := groupSet[gID]; ok { - return true - } - } - for _, gID := range rule.Destinations { - if _, ok := groupSet[gID]; ok { - return true - } - } - } - return false -} - -func routeReferencesGroups(r *route.Route, groupSet map[string]struct{}) bool { - for _, gID := range r.Groups { - if _, ok := groupSet[gID]; ok { - return true - } - } - for _, gID := range r.PeerGroups { - if _, ok := groupSet[gID]; ok { - return true - } - } - for _, gID := range r.AccessControlGroups { - if _, ok := groupSet[gID]; ok { - return true - } - } - return false -} - -func routerReferencesGroups(router *routerTypes.NetworkRouter, groupSet map[string]struct{}) bool { - for _, gID := range router.PeerGroups { - if _, ok := groupSet[gID]; ok { - return true - } - } - return false -} diff --git a/management/server/networks/manager.go b/management/server/networks/manager.go index 0f21ea9ba..7191694ac 100644 --- a/management/server/networks/manager.go +++ b/management/server/networks/manager.go @@ -245,62 +245,84 @@ func resolveNetworkAffectedPeers(ctx context.Context, s store.Store, accountID s } if len(data.resourceGroupIDs) > 0 { - destSet := make(map[string]struct{}, len(data.resourceGroupIDs)) for _, gID := range data.resourceGroupIDs { - destSet[gID] = struct{}{} groupSet[gID] = struct{}{} } - - for _, policy := range data.policies { - if policy == nil || !policy.Enabled { - continue - } - for _, rule := range policy.Rules { - if rule == nil || !rule.Enabled { - continue - } - for _, gID := range rule.Destinations { - if _, ok := destSet[gID]; ok { - for _, srcGID := range rule.Sources { - groupSet[srcGID] = struct{}{} - } - break - } - } - } - } + collectPolicySourceGroups(data.policies, data.resourceGroupIDs, groupSet) } if len(groupSet) == 0 && len(data.directPeerIDs) == 0 { return nil } + peerIDs := resolveGroupsAndDirectPeers(ctx, s, accountID, groupSet, data.directPeerIDs) + + log.WithContext(ctx).Tracef("resolveNetworkAffectedPeers: result %d peers: %v", len(peerIDs), peerIDs) + return peerIDs +} + +// collectPolicySourceGroups finds policies whose rules reference any of the destination group IDs +// and adds their source groups to the groupSet. +func collectPolicySourceGroups(policies []*nbTypes.Policy, destGroupIDs []string, groupSet map[string]struct{}) { + destSet := make(map[string]struct{}, len(destGroupIDs)) + for _, gID := range destGroupIDs { + destSet[gID] = struct{}{} + } + + for _, policy := range policies { + if policy == nil || !policy.Enabled { + continue + } + for _, rule := range policy.Rules { + if rule == nil || !rule.Enabled { + continue + } + if ruleMatchesDestinations(rule, destSet) { + for _, gID := range rule.Sources { + groupSet[gID] = struct{}{} + } + } + } + } +} + +// ruleMatchesDestinations checks if a policy rule references any of the destination groups. +func ruleMatchesDestinations(rule *nbTypes.PolicyRule, destSet map[string]struct{}) bool { + for _, gID := range rule.Destinations { + if _, ok := destSet[gID]; ok { + return true + } + } + return false +} + +// resolveGroupsAndDirectPeers resolves group IDs and direct peer IDs into a deduplicated peer ID list. +func resolveGroupsAndDirectPeers(ctx context.Context, s store.Store, accountID string, groupSet map[string]struct{}, directPeerIDs []string) []string { groupIDs := make([]string, 0, len(groupSet)) for gID := range groupSet { groupIDs = append(groupIDs, gID) } - log.WithContext(ctx).Tracef("resolveNetworkAffectedPeers: resolved groupIDs=%v", groupIDs) peerIDs, err := s.GetPeerIDsByGroups(ctx, accountID, groupIDs) if err != nil { log.WithContext(ctx).Errorf("failed to resolve peer IDs: %v", err) return nil } - if len(data.directPeerIDs) > 0 { - seen := make(map[string]struct{}, len(peerIDs)) - for _, id := range peerIDs { - seen[id] = struct{}{} - } - for _, id := range data.directPeerIDs { - if _, exists := seen[id]; !exists { - peerIDs = append(peerIDs, id) - seen[id] = struct{}{} - } - } + if len(directPeerIDs) == 0 { + return peerIDs } - log.WithContext(ctx).Tracef("resolveNetworkAffectedPeers: result %d peers: %v", len(peerIDs), peerIDs) + seen := make(map[string]struct{}, len(peerIDs)) + for _, id := range peerIDs { + seen[id] = struct{}{} + } + for _, id := range directPeerIDs { + if _, exists := seen[id]; !exists { + peerIDs = append(peerIDs, id) + seen[id] = struct{}{} + } + } return peerIDs } diff --git a/management/server/networks/resources/manager.go b/management/server/networks/resources/manager.go index 03dbe542b..552daf37f 100644 --- a/management/server/networks/resources/manager.go +++ b/management/server/networks/resources/manager.go @@ -116,49 +116,9 @@ func (m *managerImpl) CreateResource(ctx context.Context, userID string, resourc var eventsToStore []func() var affectedData *resourceAffectedPeersData err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { - _, err = transaction.GetNetworkResourceByName(ctx, store.LockingStrengthNone, resource.AccountID, resource.Name) - if err == nil { - return status.Errorf(status.InvalidArgument, "resource with name %s already exists", resource.Name) - } - - network, err := transaction.GetNetworkByID(ctx, store.LockingStrengthUpdate, resource.AccountID, resource.NetworkID) - if err != nil { - return fmt.Errorf("failed to get network: %w", err) - } - - err = transaction.SaveNetworkResource(ctx, resource) - if err != nil { - return fmt.Errorf("failed to save network resource: %w", err) - } - - event := func() { - m.accountManager.StoreEvent(ctx, userID, resource.ID, resource.AccountID, activity.NetworkResourceCreated, resource.EventMeta(network)) - } - eventsToStore = append(eventsToStore, event) - - res := nbtypes.Resource{ - ID: resource.ID, - Type: nbtypes.ResourceType(resource.Type.String()), - } - for _, groupID := range resource.GroupIDs { - event, err := m.groupsManager.AddResourceToGroupInTransaction(ctx, transaction, resource.AccountID, userID, groupID, &res) - if err != nil { - return fmt.Errorf("failed to add resource to group: %w", err) - } - eventsToStore = append(eventsToStore, event) - } - - err = transaction.IncrementNetworkSerial(ctx, resource.AccountID) - if err != nil { - return fmt.Errorf("failed to increment network serial: %w", err) - } - - affectedData, err = loadResourceAffectedPeersData(ctx, transaction, resource.AccountID, resource.NetworkID, resource.GroupIDs) - if err != nil { - log.WithContext(ctx).Errorf("failed to load affected peers data: %v", err) - } - - return nil + var txErr error + eventsToStore, affectedData, txErr = m.createResourceInTransaction(ctx, transaction, userID, resource) + return txErr }) if err != nil { return nil, fmt.Errorf("failed to create network resource: %w", err) @@ -178,6 +138,50 @@ func (m *managerImpl) CreateResource(ctx context.Context, userID string, resourc return resource, nil } +func (m *managerImpl) createResourceInTransaction(ctx context.Context, transaction store.Store, userID string, resource *types.NetworkResource) ([]func(), *resourceAffectedPeersData, error) { + _, err := transaction.GetNetworkResourceByName(ctx, store.LockingStrengthNone, resource.AccountID, resource.Name) + if err == nil { + return nil, nil, status.Errorf(status.InvalidArgument, "resource with name %s already exists", resource.Name) + } + + network, err := transaction.GetNetworkByID(ctx, store.LockingStrengthUpdate, resource.AccountID, resource.NetworkID) + if err != nil { + return nil, nil, fmt.Errorf("failed to get network: %w", err) + } + + if err = transaction.SaveNetworkResource(ctx, resource); err != nil { + return nil, nil, fmt.Errorf("failed to save network resource: %w", err) + } + + var eventsToStore []func() + eventsToStore = append(eventsToStore, func() { + m.accountManager.StoreEvent(ctx, userID, resource.ID, resource.AccountID, activity.NetworkResourceCreated, resource.EventMeta(network)) + }) + + res := nbtypes.Resource{ + ID: resource.ID, + Type: nbtypes.ResourceType(resource.Type.String()), + } + for _, groupID := range resource.GroupIDs { + event, err := m.groupsManager.AddResourceToGroupInTransaction(ctx, transaction, resource.AccountID, userID, groupID, &res) + if err != nil { + return nil, nil, fmt.Errorf("failed to add resource to group: %w", err) + } + eventsToStore = append(eventsToStore, event) + } + + if err = transaction.IncrementNetworkSerial(ctx, resource.AccountID); err != nil { + return nil, nil, fmt.Errorf("failed to increment network serial: %w", err) + } + + affectedData, err := loadResourceAffectedPeersData(ctx, transaction, resource.AccountID, resource.NetworkID, resource.GroupIDs) + if err != nil { + log.WithContext(ctx).Errorf("failed to load affected peers data: %v", err) + } + + return eventsToStore, affectedData, nil +} + func (m *managerImpl) GetResource(ctx context.Context, accountID, userID, networkID, resourceID string) (*types.NetworkResource, error) { ok, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.Networks, operations.Read) if err != nil { @@ -502,40 +506,9 @@ func (m *managerImpl) resolveResourceAffectedPeers(ctx context.Context, accountI log.WithContext(ctx).Tracef("resolveResourceAffectedPeers: resourceGroupIDs=%v, routerPeerGroups=%v, routerDirectPeers=%v, policies=%d", data.resourceGroupIDs, data.routerPeerGroups, data.routerDirectPeers, len(data.policies)) + groupSet := make(map[string]struct{}) - var directPeerIDs []string - - destSet := make(map[string]struct{}, len(data.resourceGroupIDs)) - for _, gID := range data.resourceGroupIDs { - destSet[gID] = struct{}{} - } - - for _, policy := range data.policies { - if policy == nil || !policy.Enabled { - continue - } - for _, rule := range policy.Rules { - if rule == nil || !rule.Enabled { - continue - } - referencesResource := false - for _, gID := range rule.Destinations { - if _, ok := destSet[gID]; ok { - referencesResource = true - break - } - } - if !referencesResource { - continue - } - for _, gID := range rule.Sources { - groupSet[gID] = struct{}{} - } - if rule.SourceResource.Type == nbtypes.ResourceTypePeer && rule.SourceResource.ID != "" { - directPeerIDs = append(directPeerIDs, rule.SourceResource.ID) - } - } - } + directPeerIDs := collectResourcePolicySourceGroups(data.policies, data.resourceGroupIDs, groupSet) for _, gID := range data.routerPeerGroups { groupSet[gID] = struct{}{} @@ -546,31 +519,78 @@ func (m *managerImpl) resolveResourceAffectedPeers(ctx context.Context, accountI return nil } + peerIDs := resolveGroupsAndDirectPeers(ctx, m.store, accountID, groupSet, directPeerIDs) + + log.WithContext(ctx).Tracef("resolveResourceAffectedPeers: result %d peers: %v", len(peerIDs), peerIDs) + return peerIDs +} + +// collectResourcePolicySourceGroups finds policies whose rules reference the resource destination groups, +// adds their source groups to groupSet, and returns any direct peer IDs from source resources. +func collectResourcePolicySourceGroups(policies []*nbtypes.Policy, destGroupIDs []string, groupSet map[string]struct{}) []string { + destSet := make(map[string]struct{}, len(destGroupIDs)) + for _, gID := range destGroupIDs { + destSet[gID] = struct{}{} + } + + var directPeerIDs []string + for _, policy := range policies { + if policy == nil || !policy.Enabled { + continue + } + for _, rule := range policy.Rules { + if rule == nil || !rule.Enabled { + continue + } + if !ruleMatchesDestinations(rule, destSet) { + continue + } + for _, gID := range rule.Sources { + groupSet[gID] = struct{}{} + } + if rule.SourceResource.Type == nbtypes.ResourceTypePeer && rule.SourceResource.ID != "" { + directPeerIDs = append(directPeerIDs, rule.SourceResource.ID) + } + } + } + return directPeerIDs +} + +func ruleMatchesDestinations(rule *nbtypes.PolicyRule, destSet map[string]struct{}) bool { + for _, gID := range rule.Destinations { + if _, ok := destSet[gID]; ok { + return true + } + } + return false +} + +func resolveGroupsAndDirectPeers(ctx context.Context, s store.Store, accountID string, groupSet map[string]struct{}, directPeerIDs []string) []string { groupIDs := make([]string, 0, len(groupSet)) for gID := range groupSet { groupIDs = append(groupIDs, gID) } - peerIDs, err := m.store.GetPeerIDsByGroups(ctx, accountID, groupIDs) + peerIDs, err := s.GetPeerIDsByGroups(ctx, accountID, groupIDs) if err != nil { log.WithContext(ctx).Errorf("failed to resolve peer IDs: %v", err) return nil } - if len(directPeerIDs) > 0 { - seen := make(map[string]struct{}, len(peerIDs)) - for _, id := range peerIDs { - seen[id] = struct{}{} - } - for _, id := range directPeerIDs { - if _, exists := seen[id]; !exists { - peerIDs = append(peerIDs, id) - seen[id] = struct{}{} - } - } + if len(directPeerIDs) == 0 { + return peerIDs } - log.WithContext(ctx).Tracef("resolveResourceAffectedPeers: result %d peers: %v", len(peerIDs), peerIDs) + seen := make(map[string]struct{}, len(peerIDs)) + for _, id := range peerIDs { + seen[id] = struct{}{} + } + for _, id := range directPeerIDs { + if _, exists := seen[id]; !exists { + peerIDs = append(peerIDs, id) + seen[id] = struct{}{} + } + } return peerIDs } diff --git a/management/server/networks/routers/manager.go b/management/server/networks/routers/manager.go index 0da184ca3..67c87fbdb 100644 --- a/management/server/networks/routers/manager.go +++ b/management/server/networks/routers/manager.go @@ -170,44 +170,9 @@ func (m *managerImpl) UpdateRouter(ctx context.Context, userID string, router *t var network *networkTypes.Network var affectedData *routerAffectedPeersData err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { - network, err = transaction.GetNetworkByID(ctx, store.LockingStrengthNone, router.AccountID, router.NetworkID) - if err != nil { - return fmt.Errorf("failed to get network: %w", err) - } - - if network.ID != router.NetworkID { - return status.NewRouterNotPartOfNetworkError(router.ID, router.NetworkID) - } - - allPeerGroups := router.PeerGroups - var directPeers []string - if router.Peer != "" { - directPeers = append(directPeers, router.Peer) - } - oldRouter, err := transaction.GetNetworkRouterByID(ctx, store.LockingStrengthNone, router.AccountID, router.ID) - if err == nil { - allPeerGroups = append(allPeerGroups, oldRouter.PeerGroups...) - if oldRouter.Peer != "" { - directPeers = append(directPeers, oldRouter.Peer) - } - } - - err = transaction.SaveNetworkRouter(ctx, router) - if err != nil { - return fmt.Errorf("failed to update network router: %w", err) - } - - err = transaction.IncrementNetworkSerial(ctx, router.AccountID) - if err != nil { - return fmt.Errorf("failed to increment network serial: %w", err) - } - - affectedData, err = loadRouterAffectedPeersData(ctx, transaction, router.AccountID, router.NetworkID, allPeerGroups, directPeers...) - if err != nil { - log.WithContext(ctx).Errorf("failed to load affected peers data: %v", err) - } - - return nil + var txErr error + network, affectedData, txErr = m.updateRouterInTransaction(ctx, transaction, router) + return txErr }) if err != nil { return nil, err @@ -225,6 +190,45 @@ func (m *managerImpl) UpdateRouter(ctx context.Context, userID string, router *t return router, nil } +func (m *managerImpl) updateRouterInTransaction(ctx context.Context, transaction store.Store, router *types.NetworkRouter) (*networkTypes.Network, *routerAffectedPeersData, error) { + network, err := transaction.GetNetworkByID(ctx, store.LockingStrengthNone, router.AccountID, router.NetworkID) + if err != nil { + return nil, nil, fmt.Errorf("failed to get network: %w", err) + } + + if network.ID != router.NetworkID { + return nil, nil, status.NewRouterNotPartOfNetworkError(router.ID, router.NetworkID) + } + + allPeerGroups := router.PeerGroups + var directPeers []string + if router.Peer != "" { + directPeers = append(directPeers, router.Peer) + } + oldRouter, err := transaction.GetNetworkRouterByID(ctx, store.LockingStrengthNone, router.AccountID, router.ID) + if err == nil { + allPeerGroups = append(allPeerGroups, oldRouter.PeerGroups...) + if oldRouter.Peer != "" { + directPeers = append(directPeers, oldRouter.Peer) + } + } + + if err = transaction.SaveNetworkRouter(ctx, router); err != nil { + return nil, nil, fmt.Errorf("failed to update network router: %w", err) + } + + if err = transaction.IncrementNetworkSerial(ctx, router.AccountID); err != nil { + return nil, nil, fmt.Errorf("failed to increment network serial: %w", err) + } + + affectedData, err := loadRouterAffectedPeersData(ctx, transaction, router.AccountID, router.NetworkID, allPeerGroups, directPeers...) + if err != nil { + log.WithContext(ctx).Errorf("failed to load affected peers data: %v", err) + } + + return network, affectedData, nil +} + func (m *managerImpl) DeleteRouter(ctx context.Context, accountID, userID, networkID, routerID string) error { ok, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.Networks, operations.Delete) if err != nil { @@ -374,65 +378,79 @@ func (m *managerImpl) resolveRouterAffectedPeers(ctx context.Context, accountID } if len(data.resourceGroupIDs) > 0 { - destSet := make(map[string]struct{}, len(data.resourceGroupIDs)) - for _, gID := range data.resourceGroupIDs { - destSet[gID] = struct{}{} - } - - for _, policy := range data.policies { - if policy == nil || !policy.Enabled { - continue - } - for _, rule := range policy.Rules { - if rule == nil || !rule.Enabled { - continue - } - referencesResource := false - for _, gID := range rule.Destinations { - if _, ok := destSet[gID]; ok { - referencesResource = true - break - } - } - if !referencesResource { - continue - } - for _, gID := range rule.Sources { - groupSet[gID] = struct{}{} - } - } - } + collectPolicySourceGroups(data.policies, data.resourceGroupIDs, groupSet) } if len(groupSet) == 0 && len(data.directPeerIDs) == 0 { return nil } + peerIDs := resolveGroupsAndDirectPeers(ctx, m.store, accountID, groupSet, data.directPeerIDs) + + log.WithContext(ctx).Tracef("resolveRouterAffectedPeers: result %d peers: %v", len(peerIDs), peerIDs) + return peerIDs +} + +// collectPolicySourceGroups finds policies whose rules reference any of the destination group IDs +// and adds their source groups to the groupSet. +func collectPolicySourceGroups(policies []*nbtypes.Policy, destGroupIDs []string, groupSet map[string]struct{}) { + destSet := make(map[string]struct{}, len(destGroupIDs)) + for _, gID := range destGroupIDs { + destSet[gID] = struct{}{} + } + + for _, policy := range policies { + if policy == nil || !policy.Enabled { + continue + } + for _, rule := range policy.Rules { + if rule == nil || !rule.Enabled { + continue + } + if ruleMatchesDestinations(rule, destSet) { + for _, gID := range rule.Sources { + groupSet[gID] = struct{}{} + } + } + } + } +} + +func ruleMatchesDestinations(rule *nbtypes.PolicyRule, destSet map[string]struct{}) bool { + for _, gID := range rule.Destinations { + if _, ok := destSet[gID]; ok { + return true + } + } + return false +} + +func resolveGroupsAndDirectPeers(ctx context.Context, s store.Store, accountID string, groupSet map[string]struct{}, directPeerIDs []string) []string { groupIDs := make([]string, 0, len(groupSet)) for gID := range groupSet { groupIDs = append(groupIDs, gID) } - peerIDs, err := m.store.GetPeerIDsByGroups(ctx, accountID, groupIDs) + peerIDs, err := s.GetPeerIDsByGroups(ctx, accountID, groupIDs) if err != nil { log.WithContext(ctx).Errorf("failed to resolve peer IDs: %v", err) return nil } - if len(data.directPeerIDs) > 0 { - seen := make(map[string]struct{}, len(peerIDs)) - for _, id := range peerIDs { - seen[id] = struct{}{} - } - for _, id := range data.directPeerIDs { - if _, exists := seen[id]; !exists { - peerIDs = append(peerIDs, id) - seen[id] = struct{}{} - } - } + if len(directPeerIDs) == 0 { + return peerIDs } - log.WithContext(ctx).Tracef("resolveRouterAffectedPeers: result %d peers: %v", len(peerIDs), peerIDs) + seen := make(map[string]struct{}, len(peerIDs)) + for _, id := range peerIDs { + seen[id] = struct{}{} + } + for _, id := range directPeerIDs { + if _, exists := seen[id]; !exists { + peerIDs = append(peerIDs, id) + seen[id] = struct{}{} + } + } return peerIDs } From 3012228b91416fd5a6e8c36cb591aa329f9319ed Mon Sep 17 00:00:00 2001 From: pascal Date: Fri, 8 May 2026 16:48:09 +0200 Subject: [PATCH 16/28] missing files --- management/server/affected_groups.go | 220 ++++++++++++++++++++++++++ management/server/group_linkage.go | 226 +++++++++++++++++++++++++++ 2 files changed, 446 insertions(+) create mode 100644 management/server/affected_groups.go create mode 100644 management/server/group_linkage.go diff --git a/management/server/affected_groups.go b/management/server/affected_groups.go new file mode 100644 index 000000000..05e40db0e --- /dev/null +++ b/management/server/affected_groups.go @@ -0,0 +1,220 @@ +package server + +import ( + "context" + + log "github.com/sirupsen/logrus" + + nbdns "github.com/netbirdio/netbird/dns" + routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/route" +) + +// collectGroupChangeAffectedGroups walks policies, routes, nameservers, DNS settings, +// and network routers to collect all group IDs and direct peer IDs affected by the changed groups. +func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedGroupIDs []string) (allGroupIDs []string, directPeerIDs []string) { + if len(changedGroupIDs) == 0 { + return nil, nil + } + + changedSet := make(map[string]struct{}, len(changedGroupIDs)) + for _, id := range changedGroupIDs { + changedSet[id] = struct{}{} + } + + log.WithContext(ctx).Tracef("collecting affected groups for changed groups %v", changedGroupIDs) + + groupSet := make(map[string]struct{}) + peerSet := make(map[string]struct{}) + + collectPolicyAffectedGroups(ctx, transaction, accountID, changedSet, groupSet, peerSet) + collectRouteAffectedGroups(ctx, transaction, accountID, changedSet, groupSet, peerSet) + collectNameServerAffectedGroups(ctx, transaction, accountID, changedSet, groupSet) + collectDNSSettingsAffectedGroups(ctx, transaction, accountID, changedSet, groupSet) + collectNetworkRouterAffectedGroups(ctx, transaction, accountID, changedSet, groupSet, peerSet) + + allGroupIDs = make([]string, 0, len(groupSet)) + for gID := range groupSet { + allGroupIDs = append(allGroupIDs, gID) + } + + directPeerIDs = make([]string, 0, len(peerSet)) + for pID := range peerSet { + directPeerIDs = append(directPeerIDs, pID) + } + + log.WithContext(ctx).Tracef("affected groups resolution: changed=%v -> affectedGroups=%v, directPeers=%v", changedGroupIDs, allGroupIDs, directPeerIDs) + + return allGroupIDs, directPeerIDs +} + +func collectPolicyAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet, peerSet map[string]struct{}) { + policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get policies for group change resolution: %v", err) + return + } + + for _, policy := range policies { + if !policyReferencesGroups(policy, changedSet) { + continue + } + ruleGroups := policy.RuleGroups() + log.WithContext(ctx).Tracef("policy %s (%s) references changed groups, adding rule groups %v", policy.ID, policy.Name, ruleGroups) + for _, gID := range ruleGroups { + groupSet[gID] = struct{}{} + } + collectPolicyDirectPeers(ctx, policy, peerSet) + } +} + +func collectPolicyDirectPeers(ctx context.Context, policy *types.Policy, peerSet map[string]struct{}) { + for _, rule := range policy.Rules { + if rule.SourceResource.Type == types.ResourceTypePeer && rule.SourceResource.ID != "" { + log.WithContext(ctx).Tracef("policy %s rule %s has direct source peer %s", policy.ID, rule.ID, rule.SourceResource.ID) + peerSet[rule.SourceResource.ID] = struct{}{} + } + if rule.DestinationResource.Type == types.ResourceTypePeer && rule.DestinationResource.ID != "" { + log.WithContext(ctx).Tracef("policy %s rule %s has direct destination peer %s", policy.ID, rule.ID, rule.DestinationResource.ID) + peerSet[rule.DestinationResource.ID] = struct{}{} + } + } +} + +func collectRouteAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet, peerSet map[string]struct{}) { + routes, err := transaction.GetAccountRoutes(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get routes for group change resolution: %v", err) + return + } + + for _, r := range routes { + if !routeReferencesGroups(r, changedSet) { + continue + } + log.WithContext(ctx).Tracef("route %s (%s) references changed groups", r.ID, r.Description) + for _, gID := range r.Groups { + groupSet[gID] = struct{}{} + } + for _, gID := range r.PeerGroups { + groupSet[gID] = struct{}{} + } + for _, gID := range r.AccessControlGroups { + groupSet[gID] = struct{}{} + } + if r.Peer != "" { + log.WithContext(ctx).Tracef("route %s has direct peer %s", r.ID, r.Peer) + peerSet[r.Peer] = struct{}{} + } + } +} + +func collectNameServerAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet map[string]struct{}) { + nsGroups, err := transaction.GetAccountNameServerGroups(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get nameserver groups for group change resolution: %v", err) + return + } + + for _, ns := range nsGroups { + if !nsReferencesGroups(ns, changedSet) { + continue + } + for _, g := range ns.Groups { + groupSet[g] = struct{}{} + } + } +} + +func nsReferencesGroups(ns *nbdns.NameServerGroup, changedSet map[string]struct{}) bool { + for _, gID := range ns.Groups { + if _, ok := changedSet[gID]; ok { + log.Tracef("nameserver group %s (%s) references changed group %s", ns.ID, ns.Name, gID) + return true + } + } + return false +} + +func collectDNSSettingsAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet map[string]struct{}) { + dnsSettings, err := transaction.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get DNS settings for group change resolution: %v", err) + return + } + + for _, gID := range dnsSettings.DisabledManagementGroups { + if _, ok := changedSet[gID]; ok { + log.WithContext(ctx).Tracef("DNS disabled management group %s matches changed group", gID) + groupSet[gID] = struct{}{} + } + } +} + +func collectNetworkRouterAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet, peerSet map[string]struct{}) { + routers, err := transaction.GetNetworkRoutersByAccountID(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get network routers for group change resolution: %v", err) + return + } + + for _, router := range routers { + if !routerReferencesGroups(router, changedSet) { + continue + } + log.WithContext(ctx).Tracef("network router %s references changed groups", router.ID) + for _, gID := range router.PeerGroups { + groupSet[gID] = struct{}{} + } + if router.Peer != "" { + log.WithContext(ctx).Tracef("network router %s has direct peer %s", router.ID, router.Peer) + peerSet[router.Peer] = struct{}{} + } + } +} + +func policyReferencesGroups(policy *types.Policy, groupSet map[string]struct{}) bool { + for _, rule := range policy.Rules { + for _, gID := range rule.Sources { + if _, ok := groupSet[gID]; ok { + return true + } + } + for _, gID := range rule.Destinations { + if _, ok := groupSet[gID]; ok { + return true + } + } + } + return false +} + +func routeReferencesGroups(r *route.Route, groupSet map[string]struct{}) bool { + for _, gID := range r.Groups { + if _, ok := groupSet[gID]; ok { + return true + } + } + for _, gID := range r.PeerGroups { + if _, ok := groupSet[gID]; ok { + return true + } + } + for _, gID := range r.AccessControlGroups { + if _, ok := groupSet[gID]; ok { + return true + } + } + return false +} + +func routerReferencesGroups(router *routerTypes.NetworkRouter, groupSet map[string]struct{}) bool { + for _, gID := range router.PeerGroups { + if _, ok := groupSet[gID]; ok { + return true + } + } + return false +} diff --git a/management/server/group_linkage.go b/management/server/group_linkage.go new file mode 100644 index 000000000..a69e7d7a7 --- /dev/null +++ b/management/server/group_linkage.go @@ -0,0 +1,226 @@ +package server + +import ( + "context" + "slices" + + log "github.com/sirupsen/logrus" + + nbdns "github.com/netbirdio/netbird/dns" + routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/route" + "github.com/netbirdio/netbird/shared/management/status" +) + +func validateDeleteGroup(ctx context.Context, transaction store.Store, group *types.Group, userID string, flowGroups []string) error { + // disable a deleting integration group if the initiator is not an admin service user + if group.Issued == types.GroupIssuedIntegration { + executingUser, err := transaction.GetUserByUserID(ctx, store.LockingStrengthNone, userID) + if err != nil { + return status.Errorf(status.Internal, "failed to get user") + } + if executingUser.Role != types.UserRoleAdmin || !executingUser.IsServiceUser { + return status.Errorf(status.PermissionDenied, "only service users with admin power can delete integration group") + } + } + + if group.IsGroupAll() { + return status.Errorf(status.InvalidArgument, "deleting group ALL is not allowed") + } + + if len(group.Resources) > 0 { + return &GroupLinkError{"network resource", group.Resources[0].ID} + } + + if slices.Contains(flowGroups, group.ID) { + return &GroupLinkError{"settings", "traffic event logging"} + } + + if isLinked, linkedRoute := isGroupLinkedToRoute(ctx, transaction, group.AccountID, group.ID); isLinked { + return &GroupLinkError{"route", string(linkedRoute.NetID)} + } + + if isLinked, linkedDns := isGroupLinkedToDns(ctx, transaction, group.AccountID, group.ID); isLinked { + return &GroupLinkError{"name server groups", linkedDns.Name} + } + + if isLinked, linkedPolicy := isGroupLinkedToPolicy(ctx, transaction, group.AccountID, group.ID); isLinked { + return &GroupLinkError{"policy", linkedPolicy.Name} + } + + if isLinked, linkedSetupKey := isGroupLinkedToSetupKey(ctx, transaction, group.AccountID, group.ID); isLinked { + return &GroupLinkError{"setup key", linkedSetupKey.Name} + } + + if isLinked, linkedUser := isGroupLinkedToUser(ctx, transaction, group.AccountID, group.ID); isLinked { + return &GroupLinkError{"user", linkedUser.Id} + } + + if isLinked, linkedRouter := isGroupLinkedToNetworkRouter(ctx, transaction, group.AccountID, group.ID); isLinked { + return &GroupLinkError{"network router", linkedRouter.ID} + } + + return checkGroupLinkedToSettings(ctx, transaction, group) +} + +// checkGroupLinkedToSettings verifies if a group is linked to any settings in the account. +func checkGroupLinkedToSettings(ctx context.Context, transaction store.Store, group *types.Group) error { + dnsSettings, err := transaction.GetAccountDNSSettings(ctx, store.LockingStrengthNone, group.AccountID) + if err != nil { + return status.Errorf(status.Internal, "failed to get DNS settings") + } + + if slices.Contains(dnsSettings.DisabledManagementGroups, group.ID) { + return &GroupLinkError{"disabled DNS management groups", group.Name} + } + + settings, err := transaction.GetAccountSettings(ctx, store.LockingStrengthNone, group.AccountID) + if err != nil { + return status.Errorf(status.Internal, "failed to get account settings") + } + + if settings.Extra != nil && slices.Contains(settings.Extra.IntegratedValidatorGroups, group.ID) { + return &GroupLinkError{"integrated validator", group.Name} + } + + return nil +} + +// isGroupLinkedToRoute checks if a group is linked to any route in the account. +func isGroupLinkedToRoute(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *route.Route) { + routes, err := transaction.GetAccountRoutes(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("error retrieving routes while checking group linkage: %v", err) + return false, nil + } + + for _, r := range routes { + isLinked := slices.Contains(r.Groups, groupID) || + slices.Contains(r.PeerGroups, groupID) || + slices.Contains(r.AccessControlGroups, groupID) + if isLinked { + return true, r + } + } + + return false, nil +} + +// isGroupLinkedToPolicy checks if a group is linked to any policy in the account. +func isGroupLinkedToPolicy(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *types.Policy) { + policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("error retrieving policies while checking group linkage: %v", err) + return false, nil + } + + for _, policy := range policies { + for _, rule := range policy.Rules { + if slices.Contains(rule.Sources, groupID) || slices.Contains(rule.Destinations, groupID) { + return true, policy + } + } + } + return false, nil +} + +// isGroupLinkedToDns checks if a group is linked to any nameserver group in the account. +func isGroupLinkedToDns(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *nbdns.NameServerGroup) { + nameServerGroups, err := transaction.GetAccountNameServerGroups(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("error retrieving name server groups while checking group linkage: %v", err) + return false, nil + } + + for _, dns := range nameServerGroups { + for _, g := range dns.Groups { + if g == groupID { + return true, dns + } + } + } + + return false, nil +} + +// isGroupLinkedToSetupKey checks if a group is linked to any setup key in the account. +func isGroupLinkedToSetupKey(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *types.SetupKey) { + setupKeys, err := transaction.GetAccountSetupKeys(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("error retrieving setup keys while checking group linkage: %v", err) + return false, nil + } + + for _, setupKey := range setupKeys { + if slices.Contains(setupKey.AutoGroups, groupID) { + return true, setupKey + } + } + return false, nil +} + +// isGroupLinkedToUser checks if a group is linked to any user in the account. +func isGroupLinkedToUser(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *types.User) { + users, err := transaction.GetAccountUsers(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("error retrieving users while checking group linkage: %v", err) + return false, nil + } + + for _, user := range users { + if slices.Contains(user.AutoGroups, groupID) { + return true, user + } + } + return false, nil +} + +// isGroupLinkedToNetworkRouter checks if a group is linked to any network router in the account. +func isGroupLinkedToNetworkRouter(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *routerTypes.NetworkRouter) { + routers, err := transaction.GetNetworkRoutersByAccountID(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("error retrieving network routers while checking group linkage: %v", err) + return false, nil + } + + for _, router := range routers { + if slices.Contains(router.PeerGroups, groupID) { + return true, router + } + } + return false, nil +} + +// areGroupChangesAffectPeers checks if any changes to the specified groups will affect peers. +func areGroupChangesAffectPeers(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) (bool, error) { + if len(groupIDs) == 0 { + return false, nil + } + + dnsSettings, err := transaction.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID) + if err != nil { + return false, err + } + + for _, groupID := range groupIDs { + if slices.Contains(dnsSettings.DisabledManagementGroups, groupID) { + return true, nil + } + if linked, _ := isGroupLinkedToDns(ctx, transaction, accountID, groupID); linked { + return true, nil + } + if linked, _ := isGroupLinkedToPolicy(ctx, transaction, accountID, groupID); linked { + return true, nil + } + if linked, _ := isGroupLinkedToRoute(ctx, transaction, accountID, groupID); linked { + return true, nil + } + if linked, _ := isGroupLinkedToNetworkRouter(ctx, transaction, accountID, groupID); linked { + return true, nil + } + } + + return false, nil +} From 1d906e411dd0a68a49feba60080cd13e5b511179 Mon Sep 17 00:00:00 2001 From: pascal Date: Fri, 8 May 2026 19:31:46 +0200 Subject: [PATCH 17/28] fix test --- management/server/store/sql_store.go | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go index ee3d76ebc..9cd7f20f0 100644 --- a/management/server/store/sql_store.go +++ b/management/server/store/sql_store.go @@ -265,7 +265,8 @@ func (s *SqlStore) AcquireGlobalLock(ctx context.Context) (unlock func()) { return unlock } -// Deprecated: Full account operations are no longer supported +// Deprecated: Full +// account operations are no longer supported func (s *SqlStore) SaveAccount(ctx context.Context, account *types.Account) error { start := time.Now() defer func() { From 5ae6c25ac003b4a83a13c66cb605bdf521349d18 Mon Sep 17 00:00:00 2001 From: pascal Date: Fri, 8 May 2026 19:31:59 +0200 Subject: [PATCH 18/28] fix test --- management/server/route_test.go | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/management/server/route_test.go b/management/server/route_test.go index 79014790f..5ae18c253 100644 --- a/management/server/route_test.go +++ b/management/server/route_test.go @@ -1962,8 +1962,10 @@ func TestRouteAccountPeersUpdate(t *testing.T) { }) - // Creating a route with no routing peer and having peers in groups should update account peers and send peer update + // Creating a route with no routing peer and having peers in groups that don't include peer1 should not send peer1 an update t.Run("creating a route with peers in PeerGroups and Groups", func(t *testing.T) { + drainPeerUpdates(updMsg) + route := route.Route{ ID: "testingRoute2", Network: netip.MustParsePrefix("192.0.2.0/32"), @@ -1979,7 +1981,7 @@ func TestRouteAccountPeersUpdate(t *testing.T) { done := make(chan struct{}) go func() { - peerShouldReceiveUpdate(t, updMsg) + peerShouldNotReceiveUpdate(t, updMsg) close(done) }() @@ -1992,8 +1994,8 @@ func TestRouteAccountPeersUpdate(t *testing.T) { select { case <-done: - case <-time.After(peerUpdateTimeout): - t.Error("timeout waiting for peerShouldReceiveUpdate") + case <-time.After(time.Second): + t.Error("timeout waiting for peerShouldNotReceiveUpdate") } }) From aa9a1a42f5a5c6b3e8ac0c3b3ccc36c84f790391 Mon Sep 17 00:00:00 2001 From: pascal Date: Fri, 8 May 2026 19:36:21 +0200 Subject: [PATCH 19/28] remove complexity --- .../server/networks/resources/manager.go | 31 +++++++++++-------- 1 file changed, 18 insertions(+), 13 deletions(-) diff --git a/management/server/networks/resources/manager.go b/management/server/networks/resources/manager.go index 552daf37f..43a04c544 100644 --- a/management/server/networks/resources/manager.go +++ b/management/server/networks/resources/manager.go @@ -538,19 +538,24 @@ func collectResourcePolicySourceGroups(policies []*nbtypes.Policy, destGroupIDs if policy == nil || !policy.Enabled { continue } - for _, rule := range policy.Rules { - if rule == nil || !rule.Enabled { - continue - } - if !ruleMatchesDestinations(rule, destSet) { - continue - } - for _, gID := range rule.Sources { - groupSet[gID] = struct{}{} - } - if rule.SourceResource.Type == nbtypes.ResourceTypePeer && rule.SourceResource.ID != "" { - directPeerIDs = append(directPeerIDs, rule.SourceResource.ID) - } + directPeerIDs = collectSourcesFromPolicyRules(policy.Rules, destSet, groupSet, directPeerIDs) + } + return directPeerIDs +} + +func collectSourcesFromPolicyRules(rules []*nbtypes.PolicyRule, destSet map[string]struct{}, groupSet map[string]struct{}, directPeerIDs []string) []string { + for _, rule := range rules { + if rule == nil || !rule.Enabled { + continue + } + if !ruleMatchesDestinations(rule, destSet) { + continue + } + for _, gID := range rule.Sources { + groupSet[gID] = struct{}{} + } + if rule.SourceResource.Type == nbtypes.ResourceTypePeer && rule.SourceResource.ID != "" { + directPeerIDs = append(directPeerIDs, rule.SourceResource.ID) } } return directPeerIDs From 6568c905c61288b17cd69cd68b2f76d9bb41abae Mon Sep 17 00:00:00 2001 From: pascal Date: Fri, 8 May 2026 19:54:14 +0200 Subject: [PATCH 20/28] fix test --- management/server/policy_test.go | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/management/server/policy_test.go b/management/server/policy_test.go index 1eae07e79..6fb573b9e 100644 --- a/management/server/policy_test.go +++ b/management/server/policy_test.go @@ -1319,12 +1319,14 @@ func TestPolicyAccountPeersUpdate(t *testing.T) { } }) - // Updating disabled policy with destination and source groups containing peers should not update account's peers - // or send peer update + // Updating disabled policy with destination and source groups containing peers should still update account's peers + // because affected peer resolution does not filter by policy enabled state t.Run("updating disabled policy with source and destination groups with peers", func(t *testing.T) { + drainPeerUpdates(updMsg) + done := make(chan struct{}) go func() { - peerShouldNotReceiveUpdate(t, updMsg) + peerShouldReceiveUpdate(t, updMsg) close(done) }() @@ -1335,8 +1337,8 @@ func TestPolicyAccountPeersUpdate(t *testing.T) { select { case <-done: - case <-time.After(time.Second): - t.Error("timeout waiting for peerShouldNotReceiveUpdate") + case <-time.After(peerUpdateTimeout): + t.Error("timeout waiting for peerShouldReceiveUpdate") } }) From 3e83164bcd9c4cb0376aa1db046c628ca9a12b59 Mon Sep 17 00:00:00 2001 From: pascal Date: Fri, 8 May 2026 20:27:47 +0200 Subject: [PATCH 21/28] fix affected group handling --- management/server/affected_groups.go | 114 +++++++++++++++++++++++++++ management/server/peer.go | 7 ++ 2 files changed, 121 insertions(+) diff --git a/management/server/affected_groups.go b/management/server/affected_groups.go index 05e40db0e..7554fd763 100644 --- a/management/server/affected_groups.go +++ b/management/server/affected_groups.go @@ -175,6 +175,120 @@ func collectNetworkRouterAffectedGroups(ctx context.Context, transaction store.S } } +// collectDirectPeerRefAffectedGroups finds entities (policies, routes, network routers) that reference +// the changed peers directly by peer ID (not via group membership) and collects the affected groups and peers. +func collectDirectPeerRefAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedPeerIDs []string) (groupIDs []string, directPeerIDs []string) { + if len(changedPeerIDs) == 0 { + return nil, nil + } + + changedSet := make(map[string]struct{}, len(changedPeerIDs)) + for _, id := range changedPeerIDs { + changedSet[id] = struct{}{} + } + + groupSet := make(map[string]struct{}) + peerSet := make(map[string]struct{}) + + collectPolicyDirectPeerRefGroups(ctx, transaction, accountID, changedSet, groupSet, peerSet) + collectRouteDirectPeerRefGroups(ctx, transaction, accountID, changedSet, groupSet, peerSet) + collectRouterDirectPeerRefGroups(ctx, transaction, accountID, changedSet, groupSet, peerSet) + + groupIDs = make([]string, 0, len(groupSet)) + for gID := range groupSet { + groupIDs = append(groupIDs, gID) + } + + directPeerIDs = make([]string, 0, len(peerSet)) + for pID := range peerSet { + directPeerIDs = append(directPeerIDs, pID) + } + + return groupIDs, directPeerIDs +} + +func collectPolicyDirectPeerRefGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet, peerSet map[string]struct{}) { + policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get policies for direct peer ref resolution: %v", err) + return + } + + for _, policy := range policies { + if !policyReferencesDirectPeers(policy, changedSet) { + continue + } + for _, gID := range policy.RuleGroups() { + groupSet[gID] = struct{}{} + } + collectPolicyDirectPeers(ctx, policy, peerSet) + } +} + +func policyReferencesDirectPeers(policy *types.Policy, changedSet map[string]struct{}) bool { + for _, rule := range policy.Rules { + if rule.SourceResource.Type == types.ResourceTypePeer && rule.SourceResource.ID != "" { + if _, ok := changedSet[rule.SourceResource.ID]; ok { + return true + } + } + if rule.DestinationResource.Type == types.ResourceTypePeer && rule.DestinationResource.ID != "" { + if _, ok := changedSet[rule.DestinationResource.ID]; ok { + return true + } + } + } + return false +} + +func collectRouteDirectPeerRefGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet, peerSet map[string]struct{}) { + routes, err := transaction.GetAccountRoutes(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get routes for direct peer ref resolution: %v", err) + return + } + + for _, r := range routes { + if r.Peer == "" { + continue + } + if _, ok := changedSet[r.Peer]; !ok { + continue + } + for _, gID := range r.Groups { + groupSet[gID] = struct{}{} + } + for _, gID := range r.PeerGroups { + groupSet[gID] = struct{}{} + } + for _, gID := range r.AccessControlGroups { + groupSet[gID] = struct{}{} + } + peerSet[r.Peer] = struct{}{} + } +} + +func collectRouterDirectPeerRefGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet, peerSet map[string]struct{}) { + routers, err := transaction.GetNetworkRoutersByAccountID(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get network routers for direct peer ref resolution: %v", err) + return + } + + for _, router := range routers { + if router.Peer == "" { + continue + } + if _, ok := changedSet[router.Peer]; !ok { + continue + } + for _, gID := range router.PeerGroups { + groupSet[gID] = struct{}{} + } + peerSet[router.Peer] = struct{}{} + } +} + func policyReferencesGroups(policy *types.Policy, groupSet map[string]struct{}) bool { for _, rule := range policy.Rules { for _, gID := range rule.Sources { diff --git a/management/server/peer.go b/management/server/peer.go index 5ea0197bc..8d25754a1 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -1344,6 +1344,13 @@ func (am *DefaultAccountManager) resolveAffectedPeersForPeerChanges(ctx context. log.WithContext(ctx).Tracef("resolveAffectedPeersForPeerChanges: changedPeers=%v -> groups=%v", changedPeerIDs, groupIDs) allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, s, accountID, groupIDs) + + // Also collect groups/peers from entities that reference the changed peers directly by ID + // (e.g. Route.Peer, PolicyRule.SourceResource/DestinationResource, NetworkRouter.Peer) + directRefGroups, directRefPeers := collectDirectPeerRefAffectedGroups(ctx, s, accountID, changedPeerIDs) + allGroupIDs = append(allGroupIDs, directRefGroups...) + directPeerIDs = append(directPeerIDs, directRefPeers...) + result := am.resolvePeerIDs(ctx, s, accountID, allGroupIDs, directPeerIDs) log.WithContext(ctx).Tracef("resolveAffectedPeersForPeerChanges: changedPeers=%v -> %d affected peers", changedPeerIDs, len(result)) From 13d26106f8b1275433ab66431c2d3bded0e62498 Mon Sep 17 00:00:00 2001 From: pascal Date: Fri, 8 May 2026 20:44:17 +0200 Subject: [PATCH 22/28] improve db calls --- management/server/affected_groups.go | 87 +++++++-------------- management/server/group_linkage.go | 108 ++++++++++++++++++++++----- 2 files changed, 114 insertions(+), 81 deletions(-) diff --git a/management/server/affected_groups.go b/management/server/affected_groups.go index 7554fd763..d0857ddec 100644 --- a/management/server/affected_groups.go +++ b/management/server/affected_groups.go @@ -66,18 +66,16 @@ func collectPolicyAffectedGroups(ctx context.Context, transaction store.Store, a for _, gID := range ruleGroups { groupSet[gID] = struct{}{} } - collectPolicyDirectPeers(ctx, policy, peerSet) + collectPolicyDirectPeers(policy, peerSet) } } -func collectPolicyDirectPeers(ctx context.Context, policy *types.Policy, peerSet map[string]struct{}) { +func collectPolicyDirectPeers(policy *types.Policy, peerSet map[string]struct{}) { for _, rule := range policy.Rules { if rule.SourceResource.Type == types.ResourceTypePeer && rule.SourceResource.ID != "" { - log.WithContext(ctx).Tracef("policy %s rule %s has direct source peer %s", policy.ID, rule.ID, rule.SourceResource.ID) peerSet[rule.SourceResource.ID] = struct{}{} } if rule.DestinationResource.Type == types.ResourceTypePeer && rule.DestinationResource.ID != "" { - log.WithContext(ctx).Tracef("policy %s rule %s has direct destination peer %s", policy.ID, rule.ID, rule.DestinationResource.ID) peerSet[rule.DestinationResource.ID] = struct{}{} } } @@ -95,17 +93,8 @@ func collectRouteAffectedGroups(ctx context.Context, transaction store.Store, ac continue } log.WithContext(ctx).Tracef("route %s (%s) references changed groups", r.ID, r.Description) - for _, gID := range r.Groups { - groupSet[gID] = struct{}{} - } - for _, gID := range r.PeerGroups { - groupSet[gID] = struct{}{} - } - for _, gID := range r.AccessControlGroups { - groupSet[gID] = struct{}{} - } + addAllToSet(groupSet, r.Groups, r.PeerGroups, r.AccessControlGroups) if r.Peer != "" { - log.WithContext(ctx).Tracef("route %s has direct peer %s", r.ID, r.Peer) peerSet[r.Peer] = struct{}{} } } @@ -221,26 +210,27 @@ func collectPolicyDirectPeerRefGroups(ctx context.Context, transaction store.Sto for _, gID := range policy.RuleGroups() { groupSet[gID] = struct{}{} } - collectPolicyDirectPeers(ctx, policy, peerSet) + collectPolicyDirectPeers(policy, peerSet) } } func policyReferencesDirectPeers(policy *types.Policy, changedSet map[string]struct{}) bool { for _, rule := range policy.Rules { - if rule.SourceResource.Type == types.ResourceTypePeer && rule.SourceResource.ID != "" { - if _, ok := changedSet[rule.SourceResource.ID]; ok { - return true - } - } - if rule.DestinationResource.Type == types.ResourceTypePeer && rule.DestinationResource.ID != "" { - if _, ok := changedSet[rule.DestinationResource.ID]; ok { - return true - } + if isDirectPeerInSet(rule.SourceResource, changedSet) || isDirectPeerInSet(rule.DestinationResource, changedSet) { + return true } } return false } +func isDirectPeerInSet(res types.Resource, set map[string]struct{}) bool { + if res.Type != types.ResourceTypePeer || res.ID == "" { + return false + } + _, ok := set[res.ID] + return ok +} + func collectRouteDirectPeerRefGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet, peerSet map[string]struct{}) { routes, err := transaction.GetAccountRoutes(ctx, store.LockingStrengthNone, accountID) if err != nil { @@ -255,15 +245,7 @@ func collectRouteDirectPeerRefGroups(ctx context.Context, transaction store.Stor if _, ok := changedSet[r.Peer]; !ok { continue } - for _, gID := range r.Groups { - groupSet[gID] = struct{}{} - } - for _, gID := range r.PeerGroups { - groupSet[gID] = struct{}{} - } - for _, gID := range r.AccessControlGroups { - groupSet[gID] = struct{}{} - } + addAllToSet(groupSet, r.Groups, r.PeerGroups, r.AccessControlGroups) peerSet[r.Peer] = struct{}{} } } @@ -291,44 +273,25 @@ func collectRouterDirectPeerRefGroups(ctx context.Context, transaction store.Sto func policyReferencesGroups(policy *types.Policy, groupSet map[string]struct{}) bool { for _, rule := range policy.Rules { - for _, gID := range rule.Sources { - if _, ok := groupSet[gID]; ok { - return true - } - } - for _, gID := range rule.Destinations { - if _, ok := groupSet[gID]; ok { - return true - } + if anyInSet(rule.Sources, groupSet) || anyInSet(rule.Destinations, groupSet) { + return true } } return false } func routeReferencesGroups(r *route.Route, groupSet map[string]struct{}) bool { - for _, gID := range r.Groups { - if _, ok := groupSet[gID]; ok { - return true - } - } - for _, gID := range r.PeerGroups { - if _, ok := groupSet[gID]; ok { - return true - } - } - for _, gID := range r.AccessControlGroups { - if _, ok := groupSet[gID]; ok { - return true - } - } - return false + return anyInSet(r.Groups, groupSet) || anyInSet(r.PeerGroups, groupSet) || anyInSet(r.AccessControlGroups, groupSet) } func routerReferencesGroups(router *routerTypes.NetworkRouter, groupSet map[string]struct{}) bool { - for _, gID := range router.PeerGroups { - if _, ok := groupSet[gID]; ok { - return true + return anyInSet(router.PeerGroups, groupSet) +} + +func addAllToSet(set map[string]struct{}, slices ...[]string) { + for _, s := range slices { + for _, id := range s { + set[id] = struct{}{} } } - return false } diff --git a/management/server/group_linkage.go b/management/server/group_linkage.go index a69e7d7a7..56491f9df 100644 --- a/management/server/group_linkage.go +++ b/management/server/group_linkage.go @@ -194,33 +194,103 @@ func isGroupLinkedToNetworkRouter(ctx context.Context, transaction store.Store, } // areGroupChangesAffectPeers checks if any changes to the specified groups will affect peers. +// It fetches each collection once and checks all groupIDs against them in memory. func areGroupChangesAffectPeers(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) (bool, error) { if len(groupIDs) == 0 { return false, nil } - dnsSettings, err := transaction.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID) - if err != nil { - return false, err + groupSet := make(map[string]struct{}, len(groupIDs)) + for _, id := range groupIDs { + groupSet[id] = struct{}{} } - for _, groupID := range groupIDs { - if slices.Contains(dnsSettings.DisabledManagementGroups, groupID) { - return true, nil - } - if linked, _ := isGroupLinkedToDns(ctx, transaction, accountID, groupID); linked { - return true, nil - } - if linked, _ := isGroupLinkedToPolicy(ctx, transaction, accountID, groupID); linked { - return true, nil - } - if linked, _ := isGroupLinkedToRoute(ctx, transaction, accountID, groupID); linked { - return true, nil - } - if linked, _ := isGroupLinkedToNetworkRouter(ctx, transaction, accountID, groupID); linked { - return true, nil - } + if affected, err := dnsSettingsReferenceGroups(ctx, transaction, accountID, groupSet); affected || err != nil { + return affected, err + } + if affected, err := nameServersReferenceGroups(ctx, transaction, accountID, groupSet); affected || err != nil { + return affected, err + } + if affected, err := policiesReferenceGroups(ctx, transaction, accountID, groupSet); affected || err != nil { + return affected, err + } + if affected, err := routesReferenceGroups(ctx, transaction, accountID, groupSet); affected || err != nil { + return affected, err + } + if affected, err := networkRoutersReferenceGroups(ctx, transaction, accountID, groupSet); affected || err != nil { + return affected, err } return false, nil } + +func anyInSet(ids []string, set map[string]struct{}) bool { + for _, id := range ids { + if _, ok := set[id]; ok { + return true + } + } + return false +} + +func dnsSettingsReferenceGroups(ctx context.Context, transaction store.Store, accountID string, groupSet map[string]struct{}) (bool, error) { + dnsSettings, err := transaction.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID) + if err != nil { + return false, err + } + return anyInSet(dnsSettings.DisabledManagementGroups, groupSet), nil +} + +func nameServersReferenceGroups(ctx context.Context, transaction store.Store, accountID string, groupSet map[string]struct{}) (bool, error) { + nameServerGroups, err := transaction.GetAccountNameServerGroups(ctx, store.LockingStrengthNone, accountID) + if err != nil { + return false, err + } + for _, ns := range nameServerGroups { + if anyInSet(ns.Groups, groupSet) { + return true, nil + } + } + return false, nil +} + +func policiesReferenceGroups(ctx context.Context, transaction store.Store, accountID string, groupSet map[string]struct{}) (bool, error) { + policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + if err != nil { + return false, err + } + for _, policy := range policies { + for _, rule := range policy.Rules { + if anyInSet(rule.Sources, groupSet) || anyInSet(rule.Destinations, groupSet) { + return true, nil + } + } + } + return false, nil +} + +func routesReferenceGroups(ctx context.Context, transaction store.Store, accountID string, groupSet map[string]struct{}) (bool, error) { + routes, err := transaction.GetAccountRoutes(ctx, store.LockingStrengthNone, accountID) + if err != nil { + return false, err + } + for _, r := range routes { + if anyInSet(r.Groups, groupSet) || anyInSet(r.PeerGroups, groupSet) || anyInSet(r.AccessControlGroups, groupSet) { + return true, nil + } + } + return false, nil +} + +func networkRoutersReferenceGroups(ctx context.Context, transaction store.Store, accountID string, groupSet map[string]struct{}) (bool, error) { + routers, err := transaction.GetNetworkRoutersByAccountID(ctx, store.LockingStrengthNone, accountID) + if err != nil { + return false, err + } + for _, router := range routers { + if anyInSet(router.PeerGroups, groupSet) { + return true, nil + } + } + return false, nil +} From c948d7398f567405d73e146cca95a62a21189828 Mon Sep 17 00:00:00 2001 From: pascal Date: Fri, 8 May 2026 20:51:46 +0200 Subject: [PATCH 23/28] further improve db calls --- management/server/affected_groups.go | 344 +++++++++++---------------- management/server/group_linkage.go | 9 - management/server/peer.go | 10 +- 3 files changed, 137 insertions(+), 226 deletions(-) diff --git a/management/server/affected_groups.go b/management/server/affected_groups.go index d0857ddec..4b765ec41 100644 --- a/management/server/affected_groups.go +++ b/management/server/affected_groups.go @@ -5,68 +5,137 @@ import ( log "github.com/sirupsen/logrus" - nbdns "github.com/netbirdio/netbird/dns" routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" "github.com/netbirdio/netbird/management/server/store" "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/route" ) -// collectGroupChangeAffectedGroups walks policies, routes, nameservers, DNS settings, -// and network routers to collect all group IDs and direct peer IDs affected by the changed groups. -func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedGroupIDs []string) (allGroupIDs []string, directPeerIDs []string) { - if len(changedGroupIDs) == 0 { +// collectPeerChangeAffectedGroups walks policies, routes, nameservers, DNS settings, +// and network routers to collect all group IDs and direct peer IDs affected by the +// changed groups and/or changed peers. Each collection is fetched from the store exactly once. +func collectPeerChangeAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedGroupIDs, changedPeerIDs []string) (allGroupIDs []string, directPeerIDs []string) { + if len(changedGroupIDs) == 0 && len(changedPeerIDs) == 0 { return nil, nil } - changedSet := make(map[string]struct{}, len(changedGroupIDs)) - for _, id := range changedGroupIDs { - changedSet[id] = struct{}{} - } - - log.WithContext(ctx).Tracef("collecting affected groups for changed groups %v", changedGroupIDs) + changedGroupSet := toSet(changedGroupIDs) + changedPeerSet := toSet(changedPeerIDs) groupSet := make(map[string]struct{}) peerSet := make(map[string]struct{}) - collectPolicyAffectedGroups(ctx, transaction, accountID, changedSet, groupSet, peerSet) - collectRouteAffectedGroups(ctx, transaction, accountID, changedSet, groupSet, peerSet) - collectNameServerAffectedGroups(ctx, transaction, accountID, changedSet, groupSet) - collectDNSSettingsAffectedGroups(ctx, transaction, accountID, changedSet, groupSet) - collectNetworkRouterAffectedGroups(ctx, transaction, accountID, changedSet, groupSet, peerSet) + collectAffectedFromPolicies(ctx, transaction, accountID, changedGroupSet, changedPeerSet, groupSet, peerSet) + collectAffectedFromRoutes(ctx, transaction, accountID, changedGroupSet, changedPeerSet, groupSet, peerSet) + collectAffectedFromNameServers(ctx, transaction, accountID, changedGroupSet, groupSet) + collectAffectedFromDNSSettings(ctx, transaction, accountID, changedGroupSet, groupSet) + collectAffectedFromNetworkRouters(ctx, transaction, accountID, changedGroupSet, changedPeerSet, groupSet, peerSet) - allGroupIDs = make([]string, 0, len(groupSet)) - for gID := range groupSet { - allGroupIDs = append(allGroupIDs, gID) - } + allGroupIDs = setToSlice(groupSet) + directPeerIDs = setToSlice(peerSet) - directPeerIDs = make([]string, 0, len(peerSet)) - for pID := range peerSet { - directPeerIDs = append(directPeerIDs, pID) - } - - log.WithContext(ctx).Tracef("affected groups resolution: changed=%v -> affectedGroups=%v, directPeers=%v", changedGroupIDs, allGroupIDs, directPeerIDs) + log.WithContext(ctx).Tracef("affected groups resolution: changedGroups=%v changedPeers=%v -> affectedGroups=%v, directPeers=%v", + changedGroupIDs, changedPeerIDs, allGroupIDs, directPeerIDs) return allGroupIDs, directPeerIDs } -func collectPolicyAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet, peerSet map[string]struct{}) { +// collectGroupChangeAffectedGroups is a convenience wrapper used by callers that only have changed groups. +func collectGroupChangeAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedGroupIDs []string) ([]string, []string) { + return collectPeerChangeAffectedGroups(ctx, transaction, accountID, changedGroupIDs, nil) +} + +func collectAffectedFromPolicies(ctx context.Context, transaction store.Store, accountID string, changedGroupSet, changedPeerSet map[string]struct{}, groupSet, peerSet map[string]struct{}) { policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) if err != nil { - log.WithContext(ctx).Errorf("failed to get policies for group change resolution: %v", err) + log.WithContext(ctx).Errorf("failed to get policies for affected group resolution: %v", err) return } for _, policy := range policies { - if !policyReferencesGroups(policy, changedSet) { + matchedByGroup := policyReferencesGroups(policy, changedGroupSet) + matchedByPeer := len(changedPeerSet) > 0 && policyReferencesDirectPeers(policy, changedPeerSet) + if !matchedByGroup && !matchedByPeer { continue } - ruleGroups := policy.RuleGroups() - log.WithContext(ctx).Tracef("policy %s (%s) references changed groups, adding rule groups %v", policy.ID, policy.Name, ruleGroups) - for _, gID := range ruleGroups { + addAllToSet(groupSet, policy.RuleGroups()) + collectPolicyDirectPeers(policy, peerSet) + } +} + +func collectAffectedFromRoutes(ctx context.Context, transaction store.Store, accountID string, changedGroupSet, changedPeerSet map[string]struct{}, groupSet, peerSet map[string]struct{}) { + routes, err := transaction.GetAccountRoutes(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get routes for affected group resolution: %v", err) + return + } + + for _, r := range routes { + matchedByGroup := routeReferencesGroups(r, changedGroupSet) + matchedByPeer := r.Peer != "" && len(changedPeerSet) > 0 && isInSet(r.Peer, changedPeerSet) + if !matchedByGroup && !matchedByPeer { + continue + } + addAllToSet(groupSet, r.Groups, r.PeerGroups, r.AccessControlGroups) + if r.Peer != "" { + peerSet[r.Peer] = struct{}{} + } + } +} + +func collectAffectedFromNameServers(ctx context.Context, transaction store.Store, accountID string, changedGroupSet map[string]struct{}, groupSet map[string]struct{}) { + if len(changedGroupSet) == 0 { + return + } + + nsGroups, err := transaction.GetAccountNameServerGroups(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get nameserver groups for affected group resolution: %v", err) + return + } + + for _, ns := range nsGroups { + if anyInSet(ns.Groups, changedGroupSet) { + addAllToSet(groupSet, ns.Groups) + } + } +} + +func collectAffectedFromDNSSettings(ctx context.Context, transaction store.Store, accountID string, changedGroupSet map[string]struct{}, groupSet map[string]struct{}) { + if len(changedGroupSet) == 0 { + return + } + + dnsSettings, err := transaction.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get DNS settings for affected group resolution: %v", err) + return + } + + for _, gID := range dnsSettings.DisabledManagementGroups { + if _, ok := changedGroupSet[gID]; ok { groupSet[gID] = struct{}{} } - collectPolicyDirectPeers(policy, peerSet) + } +} + +func collectAffectedFromNetworkRouters(ctx context.Context, transaction store.Store, accountID string, changedGroupSet, changedPeerSet map[string]struct{}, groupSet, peerSet map[string]struct{}) { + routers, err := transaction.GetNetworkRoutersByAccountID(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("failed to get network routers for affected group resolution: %v", err) + return + } + + for _, router := range routers { + matchedByGroup := routerReferencesGroups(router, changedGroupSet) + matchedByPeer := router.Peer != "" && len(changedPeerSet) > 0 && isInSet(router.Peer, changedPeerSet) + if !matchedByGroup && !matchedByPeer { + continue + } + addAllToSet(groupSet, router.PeerGroups) + if router.Peer != "" { + peerSet[router.Peer] = struct{}{} + } } } @@ -81,139 +150,15 @@ func collectPolicyDirectPeers(policy *types.Policy, peerSet map[string]struct{}) } } -func collectRouteAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet, peerSet map[string]struct{}) { - routes, err := transaction.GetAccountRoutes(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get routes for group change resolution: %v", err) - return - } - - for _, r := range routes { - if !routeReferencesGroups(r, changedSet) { - continue - } - log.WithContext(ctx).Tracef("route %s (%s) references changed groups", r.ID, r.Description) - addAllToSet(groupSet, r.Groups, r.PeerGroups, r.AccessControlGroups) - if r.Peer != "" { - peerSet[r.Peer] = struct{}{} - } - } -} - -func collectNameServerAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet map[string]struct{}) { - nsGroups, err := transaction.GetAccountNameServerGroups(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get nameserver groups for group change resolution: %v", err) - return - } - - for _, ns := range nsGroups { - if !nsReferencesGroups(ns, changedSet) { - continue - } - for _, g := range ns.Groups { - groupSet[g] = struct{}{} - } - } -} - -func nsReferencesGroups(ns *nbdns.NameServerGroup, changedSet map[string]struct{}) bool { - for _, gID := range ns.Groups { - if _, ok := changedSet[gID]; ok { - log.Tracef("nameserver group %s (%s) references changed group %s", ns.ID, ns.Name, gID) +func policyReferencesGroups(policy *types.Policy, groupSet map[string]struct{}) bool { + for _, rule := range policy.Rules { + if anyInSet(rule.Sources, groupSet) || anyInSet(rule.Destinations, groupSet) { return true } } return false } -func collectDNSSettingsAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet map[string]struct{}) { - dnsSettings, err := transaction.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get DNS settings for group change resolution: %v", err) - return - } - - for _, gID := range dnsSettings.DisabledManagementGroups { - if _, ok := changedSet[gID]; ok { - log.WithContext(ctx).Tracef("DNS disabled management group %s matches changed group", gID) - groupSet[gID] = struct{}{} - } - } -} - -func collectNetworkRouterAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet, peerSet map[string]struct{}) { - routers, err := transaction.GetNetworkRoutersByAccountID(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get network routers for group change resolution: %v", err) - return - } - - for _, router := range routers { - if !routerReferencesGroups(router, changedSet) { - continue - } - log.WithContext(ctx).Tracef("network router %s references changed groups", router.ID) - for _, gID := range router.PeerGroups { - groupSet[gID] = struct{}{} - } - if router.Peer != "" { - log.WithContext(ctx).Tracef("network router %s has direct peer %s", router.ID, router.Peer) - peerSet[router.Peer] = struct{}{} - } - } -} - -// collectDirectPeerRefAffectedGroups finds entities (policies, routes, network routers) that reference -// the changed peers directly by peer ID (not via group membership) and collects the affected groups and peers. -func collectDirectPeerRefAffectedGroups(ctx context.Context, transaction store.Store, accountID string, changedPeerIDs []string) (groupIDs []string, directPeerIDs []string) { - if len(changedPeerIDs) == 0 { - return nil, nil - } - - changedSet := make(map[string]struct{}, len(changedPeerIDs)) - for _, id := range changedPeerIDs { - changedSet[id] = struct{}{} - } - - groupSet := make(map[string]struct{}) - peerSet := make(map[string]struct{}) - - collectPolicyDirectPeerRefGroups(ctx, transaction, accountID, changedSet, groupSet, peerSet) - collectRouteDirectPeerRefGroups(ctx, transaction, accountID, changedSet, groupSet, peerSet) - collectRouterDirectPeerRefGroups(ctx, transaction, accountID, changedSet, groupSet, peerSet) - - groupIDs = make([]string, 0, len(groupSet)) - for gID := range groupSet { - groupIDs = append(groupIDs, gID) - } - - directPeerIDs = make([]string, 0, len(peerSet)) - for pID := range peerSet { - directPeerIDs = append(directPeerIDs, pID) - } - - return groupIDs, directPeerIDs -} - -func collectPolicyDirectPeerRefGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet, peerSet map[string]struct{}) { - policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get policies for direct peer ref resolution: %v", err) - return - } - - for _, policy := range policies { - if !policyReferencesDirectPeers(policy, changedSet) { - continue - } - for _, gID := range policy.RuleGroups() { - groupSet[gID] = struct{}{} - } - collectPolicyDirectPeers(policy, peerSet) - } -} - func policyReferencesDirectPeers(policy *types.Policy, changedSet map[string]struct{}) bool { for _, rule := range policy.Rules { if isDirectPeerInSet(rule.SourceResource, changedSet) || isDirectPeerInSet(rule.DestinationResource, changedSet) { @@ -231,55 +176,6 @@ func isDirectPeerInSet(res types.Resource, set map[string]struct{}) bool { return ok } -func collectRouteDirectPeerRefGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet, peerSet map[string]struct{}) { - routes, err := transaction.GetAccountRoutes(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get routes for direct peer ref resolution: %v", err) - return - } - - for _, r := range routes { - if r.Peer == "" { - continue - } - if _, ok := changedSet[r.Peer]; !ok { - continue - } - addAllToSet(groupSet, r.Groups, r.PeerGroups, r.AccessControlGroups) - peerSet[r.Peer] = struct{}{} - } -} - -func collectRouterDirectPeerRefGroups(ctx context.Context, transaction store.Store, accountID string, changedSet, groupSet, peerSet map[string]struct{}) { - routers, err := transaction.GetNetworkRoutersByAccountID(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get network routers for direct peer ref resolution: %v", err) - return - } - - for _, router := range routers { - if router.Peer == "" { - continue - } - if _, ok := changedSet[router.Peer]; !ok { - continue - } - for _, gID := range router.PeerGroups { - groupSet[gID] = struct{}{} - } - peerSet[router.Peer] = struct{}{} - } -} - -func policyReferencesGroups(policy *types.Policy, groupSet map[string]struct{}) bool { - for _, rule := range policy.Rules { - if anyInSet(rule.Sources, groupSet) || anyInSet(rule.Destinations, groupSet) { - return true - } - } - return false -} - func routeReferencesGroups(r *route.Route, groupSet map[string]struct{}) bool { return anyInSet(r.Groups, groupSet) || anyInSet(r.PeerGroups, groupSet) || anyInSet(r.AccessControlGroups, groupSet) } @@ -288,6 +184,20 @@ func routerReferencesGroups(router *routerTypes.NetworkRouter, groupSet map[stri return anyInSet(router.PeerGroups, groupSet) } +func anyInSet(ids []string, set map[string]struct{}) bool { + for _, id := range ids { + if _, ok := set[id]; ok { + return true + } + } + return false +} + +func isInSet(id string, set map[string]struct{}) bool { + _, ok := set[id] + return ok +} + func addAllToSet(set map[string]struct{}, slices ...[]string) { for _, s := range slices { for _, id := range s { @@ -295,3 +205,19 @@ func addAllToSet(set map[string]struct{}, slices ...[]string) { } } } + +func toSet(ids []string) map[string]struct{} { + set := make(map[string]struct{}, len(ids)) + for _, id := range ids { + set[id] = struct{}{} + } + return set +} + +func setToSlice(set map[string]struct{}) []string { + s := make([]string, 0, len(set)) + for id := range set { + s = append(s, id) + } + return s +} diff --git a/management/server/group_linkage.go b/management/server/group_linkage.go index 56491f9df..626ef7956 100644 --- a/management/server/group_linkage.go +++ b/management/server/group_linkage.go @@ -224,15 +224,6 @@ func areGroupChangesAffectPeers(ctx context.Context, transaction store.Store, ac return false, nil } -func anyInSet(ids []string, set map[string]struct{}) bool { - for _, id := range ids { - if _, ok := set[id]; ok { - return true - } - } - return false -} - func dnsSettingsReferenceGroups(ctx context.Context, transaction store.Store, accountID string, groupSet map[string]struct{}) (bool, error) { dnsSettings, err := transaction.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID) if err != nil { diff --git a/management/server/peer.go b/management/server/peer.go index 8d25754a1..94c015e05 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -1343,14 +1343,8 @@ func (am *DefaultAccountManager) resolveAffectedPeersForPeerChanges(ctx context. log.WithContext(ctx).Tracef("resolveAffectedPeersForPeerChanges: changedPeers=%v -> groups=%v", changedPeerIDs, groupIDs) - allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, s, accountID, groupIDs) - - // Also collect groups/peers from entities that reference the changed peers directly by ID - // (e.g. Route.Peer, PolicyRule.SourceResource/DestinationResource, NetworkRouter.Peer) - directRefGroups, directRefPeers := collectDirectPeerRefAffectedGroups(ctx, s, accountID, changedPeerIDs) - allGroupIDs = append(allGroupIDs, directRefGroups...) - directPeerIDs = append(directPeerIDs, directRefPeers...) - + // Single pass: find entities referencing the changed groups OR the changed peers directly + allGroupIDs, directPeerIDs := collectPeerChangeAffectedGroups(ctx, s, accountID, groupIDs, changedPeerIDs) result := am.resolvePeerIDs(ctx, s, accountID, allGroupIDs, directPeerIDs) log.WithContext(ctx).Tracef("resolveAffectedPeersForPeerChanges: changedPeers=%v -> %d affected peers", changedPeerIDs, len(result)) From 80966ab1b09bd86b7a526d9402b6a47438bc0943 Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Wed, 20 May 2026 08:25:30 +0200 Subject: [PATCH 24/28] [management] Ensure SessionStartedAt has a default value (#6211) * [management] Ensure SessionStartedAt has a default value Avoid null values for the new column * [management] Add PeerStatus with LastSeen in peer_test * [management] Add migration for PeerStatusSessionStartedAt default value * [management] Add PeerStatus with LastSeen in migration tests --- management/server/migration/migration_test.go | 6 +++++- management/server/peer/peer.go | 2 +- management/server/peer_test.go | 3 +++ management/server/store/store.go | 3 +++ 4 files changed, 12 insertions(+), 2 deletions(-) diff --git a/management/server/migration/migration_test.go b/management/server/migration/migration_test.go index 5e00976c2..cc97c2dff 100644 --- a/management/server/migration/migration_test.go +++ b/management/server/migration/migration_test.go @@ -198,7 +198,11 @@ func TestMigrateNetIPFieldFromBlobToJSON_WithJSONData(t *testing.T) { require.NoError(t, err, "Failed to insert account") account.PeersG = []nbpeer.Peer{ - {AccountID: "1234", Location: nbpeer.Location{ConnectionIP: net.IP{10, 0, 0, 1}}}, + { + AccountID: "1234", + Location: nbpeer.Location{ConnectionIP: net.IP{10, 0, 0, 1}}, + Status: &nbpeer.PeerStatus{LastSeen: time.Now()}, + }, } err = db.Save(account).Error diff --git a/management/server/peer/peer.go b/management/server/peer/peer.go index 2963dfcbd..6294d1c0a 100644 --- a/management/server/peer/peer.go +++ b/management/server/peer/peer.go @@ -86,7 +86,7 @@ type PeerStatus struct { //nolint:revive // active session". Integer nanoseconds are used so equality is // precision-safe across drivers, and so the predicates compose to a // single bigint comparison. - SessionStartedAt int64 + SessionStartedAt int64 `gorm:"not null;default:0"` // Connected indicates whether peer is connected to the management service or not Connected bool // LoginExpired diff --git a/management/server/peer_test.go b/management/server/peer_test.go index 07acf865f..9d6856740 100644 --- a/management/server/peer_test.go +++ b/management/server/peer_test.go @@ -2218,6 +2218,9 @@ func Test_IsUniqueConstraintError(t *testing.T) { ID: "test-peer-id", AccountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", DNSLabel: "test-peer-dns-label", + Status: &nbpeer.PeerStatus{ + LastSeen: time.Now(), + }, } for _, tt := range tests { diff --git a/management/server/store/store.go b/management/server/store/store.go index a723c1fc3..045f1576a 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -471,6 +471,9 @@ func getMigrationsPreAuto(ctx context.Context) []migrationFunc { func(db *gorm.DB) error { return migration.MigrateNewField[types.User](ctx, db, "email", "") }, + func(db *gorm.DB) error { + return migration.MigrateNewField[nbpeer.Peer](ctx, db, "peer_status_session_started_at", int64(0)) + }, func(db *gorm.DB) error { return migration.RemoveDuplicatePeerKeys(ctx, db) }, From d250f92c435bac83fd55f00fad3ee2c292eee910 Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Wed, 20 May 2026 10:08:34 +0200 Subject: [PATCH 25/28] feat(reverse-proxy): clusters API surfaces type, online status, and capability flags (#6148) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The cluster listing now answers three questions in one round-trip instead of forcing the dashboard to cross-reference the domains API: which clusters can this account see, are they currently up, and what do they support. The ProxyCluster wire type drops the boolean self_hosted in favour of a `type` enum (`account` / `shared`) plus explicit `online`, `supports_custom_ports`, `require_subdomain`, and `supports_crowdsec` fields. Store query reworked so offline clusters still appear (no last_seen WHERE), with online and connected_proxies both derived from the existing 2-min active window via portable CASE expressions; the 1-hour heartbeat reaper still removes long-stale rows. Service manager enriches each cluster with the capability flags via the existing per-cluster lookups (CapabilityProvider now also exposes ClusterSupportsCrowdSec). GetActiveClusterAddresses* keep their tight 2-min filter so service routing and domain enumeration aren't pulled into the wider window. The hard cut removes self_hosted from the response — the dashboard is the only consumer and is updated in the matching PR; no transitional field is shipped. Adds a cross-engine regression test asserting offline clusters surface, connected_proxies counts only fresh proxies, and account-scoped BYOP clusters never leak across accounts. --- .../reverseproxy/proxy/manager/manager.go | 2 +- .../proxy/manager/manager_test.go | 2 +- .../modules/reverseproxy/proxy/proxy.go | 27 ++++- .../modules/reverseproxy/service/interface.go | 2 +- .../reverseproxy/service/interface_mock.go | 72 ++++++------ .../reverseproxy/service/manager/api.go | 14 ++- .../reverseproxy/service/manager/manager.go | 22 +++- .../shared/grpc/proxy_group_access_test.go | 2 +- .../shared/grpc/validate_session_test.go | 2 +- .../proxy/auth_callback_integration_test.go | 2 +- management/server/store/sql_store.go | 64 ++++++++-- .../store/sql_store_proxy_clusters_test.go | 109 ++++++++++++++++++ management/server/store/store.go | 2 +- management/server/store/store_mock.go | 86 +++++++------- proxy/management_integration_test.go | 2 +- shared/management/http/api/openapi.yml | 32 ++++- shared/management/http/api/types.gen.go | 73 +++++++++--- 17 files changed, 393 insertions(+), 122 deletions(-) create mode 100644 management/server/store/sql_store_proxy_clusters_test.go diff --git a/management/internals/modules/reverseproxy/proxy/manager/manager.go b/management/internals/modules/reverseproxy/proxy/manager/manager.go index b72a6ebe5..510500e0c 100644 --- a/management/internals/modules/reverseproxy/proxy/manager/manager.go +++ b/management/internals/modules/reverseproxy/proxy/manager/manager.go @@ -17,7 +17,7 @@ type store interface { UpdateProxyHeartbeat(ctx context.Context, p *proxy.Proxy) error GetActiveProxyClusterAddresses(ctx context.Context) ([]string, error) GetActiveProxyClusterAddressesForAccount(ctx context.Context, accountID string) ([]string, error) - GetActiveProxyClusters(ctx context.Context, accountID string) ([]proxy.Cluster, error) + GetProxyClusters(ctx context.Context, accountID string) ([]proxy.Cluster, error) GetClusterSupportsCustomPorts(ctx context.Context, clusterAddr string) *bool GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool diff --git a/management/internals/modules/reverseproxy/proxy/manager/manager_test.go b/management/internals/modules/reverseproxy/proxy/manager/manager_test.go index 3c53fe684..3436216b4 100644 --- a/management/internals/modules/reverseproxy/proxy/manager/manager_test.go +++ b/management/internals/modules/reverseproxy/proxy/manager/manager_test.go @@ -57,7 +57,7 @@ func (m *mockStore) GetActiveProxyClusterAddressesForAccount(ctx context.Context } return nil, nil } -func (m *mockStore) GetActiveProxyClusters(_ context.Context, _ string) ([]proxy.Cluster, error) { +func (m *mockStore) GetProxyClusters(_ context.Context, _ string) ([]proxy.Cluster, error) { return nil, nil } func (m *mockStore) CleanupStaleProxies(ctx context.Context, d time.Duration) error { diff --git a/management/internals/modules/reverseproxy/proxy/proxy.go b/management/internals/modules/reverseproxy/proxy/proxy.go index 64394799e..9da7910df 100644 --- a/management/internals/modules/reverseproxy/proxy/proxy.go +++ b/management/internals/modules/reverseproxy/proxy/proxy.go @@ -42,10 +42,35 @@ func (Proxy) TableName() string { return "proxies" } +// ClusterType is the source of a proxy cluster. +type ClusterType string + +const ( + // ClusterTypeAccount is a cluster operated by the account itself (BYOP) — + // at least one proxy row in the cluster carries a non-NULL account_id. + ClusterTypeAccount ClusterType = "account" + // ClusterTypeShared is a cluster operated by NetBird and shared across + // accounts — all proxy rows in the cluster have account_id IS NULL. + ClusterTypeShared ClusterType = "shared" +) + // Cluster represents a group of proxy nodes serving the same address. +// +// Online and ConnectedProxies derive from the same 2-min active window +// the rest of the module uses, but Cluster rows are not gated on it — +// the cluster listing surfaces offline clusters too so operators can +// see and clean them up. The 1-hour heartbeat reaper still bounds the +// table eventually. type Cluster struct { ID string Address string + Type ClusterType + Online bool ConnectedProxies int - SelfHosted bool + // Capability flags. *bool because nil means "no proxy reported a + // capability for this cluster" — the dashboard renders these as + // unknown rather than false. + SupportsCustomPorts *bool + RequireSubdomain *bool + SupportsCrowdSec *bool } diff --git a/management/internals/modules/reverseproxy/service/interface.go b/management/internals/modules/reverseproxy/service/interface.go index 6a94aa32b..dddf6ae8a 100644 --- a/management/internals/modules/reverseproxy/service/interface.go +++ b/management/internals/modules/reverseproxy/service/interface.go @@ -9,7 +9,7 @@ import ( ) type Manager interface { - GetActiveClusters(ctx context.Context, accountID, userID string) ([]proxy.Cluster, error) + GetClusters(ctx context.Context, accountID, userID string) ([]proxy.Cluster, error) DeleteAccountCluster(ctx context.Context, accountID, userID, clusterAddress string) error GetAllServices(ctx context.Context, accountID, userID string) ([]*Service, error) GetService(ctx context.Context, accountID, userID, serviceID string) (*Service, error) diff --git a/management/internals/modules/reverseproxy/service/interface_mock.go b/management/internals/modules/reverseproxy/service/interface_mock.go index 83b2162ed..24963fe30 100644 --- a/management/internals/modules/reverseproxy/service/interface_mock.go +++ b/management/internals/modules/reverseproxy/service/interface_mock.go @@ -65,20 +65,6 @@ func (mr *MockManagerMockRecorder) CreateServiceFromPeer(ctx, accountID, peerID, return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CreateServiceFromPeer", reflect.TypeOf((*MockManager)(nil).CreateServiceFromPeer), ctx, accountID, peerID, req) } -// DeleteAllServices mocks base method. -func (m *MockManager) DeleteAllServices(ctx context.Context, accountID, userID string) error { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "DeleteAllServices", ctx, accountID, userID) - ret0, _ := ret[0].(error) - return ret0 -} - -// DeleteAllServices indicates an expected call of DeleteAllServices. -func (mr *MockManagerMockRecorder) DeleteAllServices(ctx, accountID, userID interface{}) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAllServices", reflect.TypeOf((*MockManager)(nil).DeleteAllServices), ctx, accountID, userID) -} - // DeleteAccountCluster mocks base method. func (m *MockManager) DeleteAccountCluster(ctx context.Context, accountID, userID, clusterAddress string) error { m.ctrl.T.Helper() @@ -93,6 +79,20 @@ func (mr *MockManagerMockRecorder) DeleteAccountCluster(ctx, accountID, userID, return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAccountCluster", reflect.TypeOf((*MockManager)(nil).DeleteAccountCluster), ctx, accountID, userID, clusterAddress) } +// DeleteAllServices mocks base method. +func (m *MockManager) DeleteAllServices(ctx context.Context, accountID, userID string) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "DeleteAllServices", ctx, accountID, userID) + ret0, _ := ret[0].(error) + return ret0 +} + +// DeleteAllServices indicates an expected call of DeleteAllServices. +func (mr *MockManagerMockRecorder) DeleteAllServices(ctx, accountID, userID interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAllServices", reflect.TypeOf((*MockManager)(nil).DeleteAllServices), ctx, accountID, userID) +} + // DeleteService mocks base method. func (m *MockManager) DeleteService(ctx context.Context, accountID, userID, serviceID string) error { m.ctrl.T.Helper() @@ -122,21 +122,6 @@ func (mr *MockManagerMockRecorder) GetAccountServices(ctx, accountID interface{} return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountServices", reflect.TypeOf((*MockManager)(nil).GetAccountServices), ctx, accountID) } -// GetActiveClusters mocks base method. -func (m *MockManager) GetActiveClusters(ctx context.Context, accountID, userID string) ([]proxy.Cluster, error) { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "GetActiveClusters", ctx, accountID, userID) - ret0, _ := ret[0].([]proxy.Cluster) - ret1, _ := ret[1].(error) - return ret0, ret1 -} - -// GetActiveClusters indicates an expected call of GetActiveClusters. -func (mr *MockManagerMockRecorder) GetActiveClusters(ctx, accountID, userID interface{}) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveClusters", reflect.TypeOf((*MockManager)(nil).GetActiveClusters), ctx, accountID, userID) -} - // GetAllServices mocks base method. func (m *MockManager) GetAllServices(ctx context.Context, accountID, userID string) ([]*Service, error) { m.ctrl.T.Helper() @@ -152,19 +137,19 @@ func (mr *MockManagerMockRecorder) GetAllServices(ctx, accountID, userID interfa return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAllServices", reflect.TypeOf((*MockManager)(nil).GetAllServices), ctx, accountID, userID) } -// GetServiceByDomain mocks base method. -func (m *MockManager) GetServiceByDomain(ctx context.Context, domain string) (*Service, error) { +// GetClusters mocks base method. +func (m *MockManager) GetClusters(ctx context.Context, accountID, userID string) ([]proxy.Cluster, error) { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "GetServiceByDomain", ctx, domain) - ret0, _ := ret[0].(*Service) + ret := m.ctrl.Call(m, "GetClusters", ctx, accountID, userID) + ret0, _ := ret[0].([]proxy.Cluster) ret1, _ := ret[1].(error) return ret0, ret1 } -// GetServiceByDomain indicates an expected call of GetServiceByDomain. -func (mr *MockManagerMockRecorder) GetServiceByDomain(ctx, domain interface{}) *gomock.Call { +// GetClusters indicates an expected call of GetClusters. +func (mr *MockManagerMockRecorder) GetClusters(ctx, accountID, userID interface{}) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetServiceByDomain", reflect.TypeOf((*MockManager)(nil).GetServiceByDomain), ctx, domain) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetClusters", reflect.TypeOf((*MockManager)(nil).GetClusters), ctx, accountID, userID) } // GetGlobalServices mocks base method. @@ -197,6 +182,21 @@ func (mr *MockManagerMockRecorder) GetService(ctx, accountID, userID, serviceID return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetService", reflect.TypeOf((*MockManager)(nil).GetService), ctx, accountID, userID, serviceID) } +// GetServiceByDomain mocks base method. +func (m *MockManager) GetServiceByDomain(ctx context.Context, domain string) (*Service, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetServiceByDomain", ctx, domain) + ret0, _ := ret[0].(*Service) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetServiceByDomain indicates an expected call of GetServiceByDomain. +func (mr *MockManagerMockRecorder) GetServiceByDomain(ctx, domain interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetServiceByDomain", reflect.TypeOf((*MockManager)(nil).GetServiceByDomain), ctx, domain) +} + // GetServiceByID mocks base method. func (m *MockManager) GetServiceByID(ctx context.Context, accountID, serviceID string) (*Service, error) { m.ctrl.T.Helper() diff --git a/management/internals/modules/reverseproxy/service/manager/api.go b/management/internals/modules/reverseproxy/service/manager/api.go index 08272077c..9d93d52ee 100644 --- a/management/internals/modules/reverseproxy/service/manager/api.go +++ b/management/internals/modules/reverseproxy/service/manager/api.go @@ -187,7 +187,7 @@ func (h *handler) getClusters(w http.ResponseWriter, r *http.Request) { return } - clusters, err := h.manager.GetActiveClusters(r.Context(), userAuth.AccountId, userAuth.UserId) + clusters, err := h.manager.GetClusters(r.Context(), userAuth.AccountId, userAuth.UserId) if err != nil { util.WriteError(r.Context(), err, w) return @@ -196,10 +196,14 @@ func (h *handler) getClusters(w http.ResponseWriter, r *http.Request) { apiClusters := make([]api.ProxyCluster, 0, len(clusters)) for _, c := range clusters { apiClusters = append(apiClusters, api.ProxyCluster{ - Id: c.ID, - Address: c.Address, - ConnectedProxies: c.ConnectedProxies, - SelfHosted: c.SelfHosted, + Id: c.ID, + Address: c.Address, + Type: api.ProxyClusterType(c.Type), + Online: c.Online, + ConnectedProxies: c.ConnectedProxies, + SupportsCustomPorts: c.SupportsCustomPorts, + RequireSubdomain: c.RequireSubdomain, + SupportsCrowdsec: c.SupportsCrowdSec, }) } diff --git a/management/internals/modules/reverseproxy/service/manager/manager.go b/management/internals/modules/reverseproxy/service/manager/manager.go index 4a8598afb..ca0c5540f 100644 --- a/management/internals/modules/reverseproxy/service/manager/manager.go +++ b/management/internals/modules/reverseproxy/service/manager/manager.go @@ -81,6 +81,7 @@ type ClusterDeriver interface { type CapabilityProvider interface { ClusterSupportsCustomPorts(ctx context.Context, clusterAddr string) *bool ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool + ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool } type Manager struct { @@ -112,8 +113,12 @@ func (m *Manager) StartExposeReaper(ctx context.Context) { m.exposeReaper.StartExposeReaper(ctx) } -// GetActiveClusters returns all active proxy clusters with their connected proxy count. -func (m *Manager) GetActiveClusters(ctx context.Context, accountID, userID string) ([]proxy.Cluster, error) { +// GetClusters returns every proxy cluster visible to the account +// (shared + its own BYOP), regardless of whether any proxy in the +// cluster is currently heartbeating. Each cluster is enriched with the +// capability flags reported by its active proxies so the dashboard can +// render feature support without a second round-trip. +func (m *Manager) GetClusters(ctx context.Context, accountID, userID string) ([]proxy.Cluster, error) { ok, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.Services, operations.Read) if err != nil { return nil, status.NewPermissionValidationError(err) @@ -122,7 +127,18 @@ func (m *Manager) GetActiveClusters(ctx context.Context, accountID, userID strin return nil, status.NewPermissionDeniedError() } - return m.store.GetActiveProxyClusters(ctx, accountID) + clusters, err := m.store.GetProxyClusters(ctx, accountID) + if err != nil { + return nil, err + } + + for i := range clusters { + clusters[i].SupportsCustomPorts = m.capabilities.ClusterSupportsCustomPorts(ctx, clusters[i].Address) + clusters[i].RequireSubdomain = m.capabilities.ClusterRequireSubdomain(ctx, clusters[i].Address) + clusters[i].SupportsCrowdSec = m.capabilities.ClusterSupportsCrowdSec(ctx, clusters[i].Address) + } + + return clusters, nil } // DeleteAccountCluster removes all proxy registrations for the given cluster address diff --git a/management/internals/shared/grpc/proxy_group_access_test.go b/management/internals/shared/grpc/proxy_group_access_test.go index 46dad5b56..5980f8a30 100644 --- a/management/internals/shared/grpc/proxy_group_access_test.go +++ b/management/internals/shared/grpc/proxy_group_access_test.go @@ -109,7 +109,7 @@ func (m *mockReverseProxyManager) GetServiceByDomain(_ context.Context, domain s return nil, errors.New("service not found for domain: " + domain) } -func (m *mockReverseProxyManager) GetActiveClusters(_ context.Context, _, _ string) ([]proxy.Cluster, error) { +func (m *mockReverseProxyManager) GetClusters(_ context.Context, _, _ string) ([]proxy.Cluster, error) { return nil, nil } diff --git a/management/internals/shared/grpc/validate_session_test.go b/management/internals/shared/grpc/validate_session_test.go index 7b7ffcfb2..774c5d1d3 100644 --- a/management/internals/shared/grpc/validate_session_test.go +++ b/management/internals/shared/grpc/validate_session_test.go @@ -322,7 +322,7 @@ func (m *testValidateSessionServiceManager) GetServiceByDomain(ctx context.Conte return m.store.GetServiceByDomain(ctx, domain) } -func (m *testValidateSessionServiceManager) GetActiveClusters(_ context.Context, _, _ string) ([]proxy.Cluster, error) { +func (m *testValidateSessionServiceManager) GetClusters(_ context.Context, _, _ string) ([]proxy.Cluster, error) { return nil, nil } diff --git a/management/server/http/handlers/proxy/auth_callback_integration_test.go b/management/server/http/handlers/proxy/auth_callback_integration_test.go index 30d8aa0e7..f08d5daf1 100644 --- a/management/server/http/handlers/proxy/auth_callback_integration_test.go +++ b/management/server/http/handlers/proxy/auth_callback_integration_test.go @@ -444,7 +444,7 @@ func (m *testServiceManager) GetServiceByDomain(ctx context.Context, domain stri return m.store.GetServiceByDomain(ctx, domain) } -func (m *testServiceManager) GetActiveClusters(_ context.Context, _, _ string) ([]nbproxy.Cluster, error) { +func (m *testServiceManager) GetClusters(_ context.Context, _, _ string) ([]nbproxy.Cluster, error) { return nil, nil } diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go index 8cf37de56..f3c6b741b 100644 --- a/management/server/store/sql_store.go +++ b/management/server/store/sql_store.go @@ -5736,19 +5736,67 @@ func (s *SqlStore) DeleteAccountCluster(ctx context.Context, clusterAddress, acc return nil } -func (s *SqlStore) GetActiveProxyClusters(ctx context.Context, accountID string) ([]proxy.Cluster, error) { - var clusters []proxy.Cluster +// GetProxyClusters returns every cluster the account can see (shared +// plus its own BYOP), regardless of whether any proxy in the cluster +// is currently heartbeating. Online and ConnectedProxies are derived +// from the 2-min active window so the dashboard can render offline +// clusters distinctly; the 1-hour heartbeat reaper still removes rows +// that go quiet for too long. +// +// AccountOwned is determined by whether any proxy row in the group +// carries a non-NULL account_id; the caller maps that to Cluster.Type. +// Capability flags are NOT filled here — the handler enriches them via +// the per-cluster capability lookups. +func (s *SqlStore) GetProxyClusters(ctx context.Context, accountID string) ([]proxy.Cluster, error) { + activeCutoff := time.Now().Add(-proxyActiveThreshold) + type clusterRow struct { + ID string + Address string + ConnectedProxies int + Online bool + AccountOwned bool + } + + var rows []clusterRow result := s.db.Model(&proxy.Proxy{}). - Select("MIN(id) as id, cluster_address as address, COUNT(*) as connected_proxies, COUNT(account_id) > 0 as self_hosted"). - Where("status = ? AND last_seen > ? AND (account_id IS NULL OR account_id = ?)", - proxy.StatusConnected, time.Now().Add(-proxyActiveThreshold), accountID). + Select( + "MIN(id) AS id, "+ + "cluster_address AS address, "+ + // COUNT(CASE WHEN ... THEN 1 END) counts only non-NULL — i.e. only + // rows that satisfy the predicate — so it works portably across + // sqlite/postgres/mysql without dialect-specific FILTER syntax. + "COUNT(CASE WHEN status = ? AND last_seen > ? THEN 1 END) AS connected_proxies, "+ + // MAX(CASE …) > 0 expresses BOOL_OR in a way Postgres tolerates + // (Postgres can't MAX a boolean column). + "MAX(CASE WHEN status = ? AND last_seen > ? THEN 1 ELSE 0 END) > 0 AS online, "+ + "MAX(CASE WHEN account_id IS NOT NULL THEN 1 ELSE 0 END) > 0 AS account_owned", + proxy.StatusConnected, activeCutoff, + proxy.StatusConnected, activeCutoff, + ). + Where("account_id IS NULL OR account_id = ?", accountID). Group("cluster_address"). - Scan(&clusters) + Scan(&rows) if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get active proxy clusters: %v", result.Error) - return nil, status.Errorf(status.Internal, "get active proxy clusters") + log.WithContext(ctx).Errorf("failed to get proxy clusters: %v", result.Error) + return nil, status.Errorf(status.Internal, "get proxy clusters") + } + + clusters := make([]proxy.Cluster, 0, len(rows)) + for _, r := range rows { + c := proxy.Cluster{ + ID: r.ID, + Address: r.Address, + Online: r.Online, + ConnectedProxies: r.ConnectedProxies, + } + if r.AccountOwned { + c.Type = proxy.ClusterTypeAccount + } else { + c.Type = proxy.ClusterTypeShared + } + clusters = append(clusters, c) } return clusters, nil diff --git a/management/server/store/sql_store_proxy_clusters_test.go b/management/server/store/sql_store_proxy_clusters_test.go new file mode 100644 index 000000000..cdacfedae --- /dev/null +++ b/management/server/store/sql_store_proxy_clusters_test.go @@ -0,0 +1,109 @@ +package store + +import ( + "context" + "os" + "runtime" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + rpproxy "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy" +) + +// TestSqlStore_GetProxyClusters_DerivesOnlineAndType guards the +// account-visible cluster list against silent regressions in two +// dimensions: +// +// 1. Online derivation: a cluster with one stale and one fresh proxy +// is online and counts only the fresh proxy; a cluster whose +// proxies all heartbeated outside the 2-min window appears offline +// with connected_proxies = 0 (rather than disappearing, which is +// what the old query did). +// 2. Type derivation: a cluster scoped to the calling account is +// reported as `account`; a cluster with account_id IS NULL is +// reported as `shared`. Clusters scoped to other accounts must not +// leak into the result. +// +// Capability flags are intentionally not asserted here — they're filled +// by the manager (handler) layer from the per-cluster capability +// lookups, not by the store query. +func TestSqlStore_GetProxyClusters_DerivesOnlineAndType(t *testing.T) { + if (os.Getenv("CI") == "true" && runtime.GOOS == "darwin") || runtime.GOOS == "windows" { + t.Skip("skip CI tests on darwin and windows") + } + + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + ctx := context.Background() + accountID := "acct-clusters" + require.NoError(t, store.SaveAccount(ctx, newAccountWithId(ctx, accountID, "user-1", ""))) + + otherAccountID := "acct-other" + require.NoError(t, store.SaveAccount(ctx, newAccountWithId(ctx, otherAccountID, "user-2", ""))) + + acctID := accountID + otherID := otherAccountID + + fresh := time.Now().Add(-30 * time.Second) + stale := time.Now().Add(-30 * time.Minute) + + mustSave := func(id, cluster string, accID *string, status string, lastSeen time.Time) { + require.NoError(t, store.SaveProxy(ctx, &rpproxy.Proxy{ + ID: id, + SessionID: id + "-sess", + ClusterAddress: cluster, + IPAddress: "10.0.0.1", + AccountID: accID, + LastSeen: lastSeen, + Status: status, + })) + } + + // shared-mixed: one fresh + one stale proxy → online, connected=1 + mustSave("p-shared-fresh", "shared-mixed.netbird.io", nil, rpproxy.StatusConnected, fresh) + mustSave("p-shared-stale", "shared-mixed.netbird.io", nil, rpproxy.StatusConnected, stale) + + // shared-offline: only stale proxies → offline, connected=0, + // but row must still appear (this is the new semantic — old + // query would have dropped it entirely). + mustSave("p-shared-off", "shared-offline.netbird.io", nil, rpproxy.StatusConnected, stale) + + // account-online: BYOP cluster owned by acctID, fresh + mustSave("p-acct-fresh", "byop.acct.example", &acctID, rpproxy.StatusConnected, fresh) + + // other-account: must not surface for acctID + mustSave("p-other", "byop.other.example", &otherID, rpproxy.StatusConnected, fresh) + + clusters, err := store.GetProxyClusters(ctx, accountID) + require.NoError(t, err) + + byAddr := map[string]rpproxy.Cluster{} + for _, c := range clusters { + byAddr[c.Address] = c + } + + assert.NotContains(t, byAddr, "byop.other.example", + "another account's BYOP cluster must not leak into this account's listing") + + require.Contains(t, byAddr, "shared-mixed.netbird.io") + mixed := byAddr["shared-mixed.netbird.io"] + assert.Equal(t, rpproxy.ClusterTypeShared, mixed.Type, "shared cluster (account_id IS NULL) must be reported as Type=shared") + assert.True(t, mixed.Online, "cluster with a fresh proxy must be online") + assert.Equal(t, 1, mixed.ConnectedProxies, "connected_proxies must count only fresh proxies; the stale one should not bump the count") + + require.Contains(t, byAddr, "shared-offline.netbird.io", + "offline clusters must still appear so the dashboard can render them — the old GetActiveProxyClusters would have dropped this row, which is the regression this test guards against") + offline := byAddr["shared-offline.netbird.io"] + assert.Equal(t, rpproxy.ClusterTypeShared, offline.Type) + assert.False(t, offline.Online, "no fresh heartbeat → offline") + assert.Equal(t, 0, offline.ConnectedProxies, "no fresh proxies → connected_proxies=0") + + require.Contains(t, byAddr, "byop.acct.example") + acct := byAddr["byop.acct.example"] + assert.Equal(t, rpproxy.ClusterTypeAccount, acct.Type, "BYOP cluster owned by the account must be reported as Type=account") + assert.True(t, acct.Online) + assert.Equal(t, 1, acct.ConnectedProxies) + }) +} diff --git a/management/server/store/store.go b/management/server/store/store.go index 045f1576a..42cdcf36d 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -307,7 +307,7 @@ type Store interface { UpdateProxyHeartbeat(ctx context.Context, p *proxy.Proxy) error GetActiveProxyClusterAddresses(ctx context.Context) ([]string, error) GetActiveProxyClusterAddressesForAccount(ctx context.Context, accountID string) ([]string, error) - GetActiveProxyClusters(ctx context.Context, accountID string) ([]proxy.Cluster, error) + GetProxyClusters(ctx context.Context, accountID string) ([]proxy.Cluster, error) GetClusterSupportsCustomPorts(ctx context.Context, clusterAddr string) *bool GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool diff --git a/management/server/store/store_mock.go b/management/server/store/store_mock.go index d51629606..4f9d875d2 100644 --- a/management/server/store/store_mock.go +++ b/management/server/store/store_mock.go @@ -380,6 +380,20 @@ func (mr *MockStoreMockRecorder) DeleteAccount(ctx, account interface{}) *gomock return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAccount", reflect.TypeOf((*MockStore)(nil).DeleteAccount), ctx, account) } +// DeleteAccountCluster mocks base method. +func (m *MockStore) DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "DeleteAccountCluster", ctx, clusterAddress, accountID) + ret0, _ := ret[0].(error) + return ret0 +} + +// DeleteAccountCluster indicates an expected call of DeleteAccountCluster. +func (mr *MockStoreMockRecorder) DeleteAccountCluster(ctx, clusterAddress, accountID interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAccountCluster", reflect.TypeOf((*MockStore)(nil).DeleteAccountCluster), ctx, clusterAddress, accountID) +} + // DeleteCustomDomain mocks base method. func (m *MockStore) DeleteCustomDomain(ctx context.Context, accountID, domainID string) error { m.ctrl.T.Helper() @@ -577,20 +591,6 @@ func (mr *MockStoreMockRecorder) DeletePostureChecks(ctx, accountID, postureChec return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeletePostureChecks", reflect.TypeOf((*MockStore)(nil).DeletePostureChecks), ctx, accountID, postureChecksID) } -// DeleteAccountCluster mocks base method. -func (m *MockStore) DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "DeleteAccountCluster", ctx, clusterAddress, accountID) - ret0, _ := ret[0].(error) - return ret0 -} - -// DeleteAccountCluster indicates an expected call of DeleteAccountCluster. -func (mr *MockStoreMockRecorder) DeleteAccountCluster(ctx, clusterAddress, accountID interface{}) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAccountCluster", reflect.TypeOf((*MockStore)(nil).DeleteAccountCluster), ctx, clusterAddress, accountID) -} - // DeleteRoute mocks base method. func (m *MockStore) DeleteRoute(ctx context.Context, accountID, routeID string) error { m.ctrl.T.Helper() @@ -731,6 +731,20 @@ func (mr *MockStoreMockRecorder) DeleteZoneDNSRecords(ctx, accountID, zoneID int return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteZoneDNSRecords", reflect.TypeOf((*MockStore)(nil).DeleteZoneDNSRecords), ctx, accountID, zoneID) } +// DisconnectProxy mocks base method. +func (m *MockStore) DisconnectProxy(ctx context.Context, proxyID, sessionID string) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "DisconnectProxy", ctx, proxyID, sessionID) + ret0, _ := ret[0].(error) + return ret0 +} + +// DisconnectProxy indicates an expected call of DisconnectProxy. +func (mr *MockStoreMockRecorder) DisconnectProxy(ctx, proxyID, sessionID interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DisconnectProxy", reflect.TypeOf((*MockStore)(nil).DisconnectProxy), ctx, proxyID, sessionID) +} + // EphemeralServiceExists mocks base method. func (m *MockStore) EphemeralServiceExists(ctx context.Context, lockStrength LockingStrength, accountID, peerID, domain string) (bool, error) { m.ctrl.T.Helper() @@ -1332,21 +1346,6 @@ func (mr *MockStoreMockRecorder) GetActiveProxyClusterAddressesForAccount(ctx, a return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveProxyClusterAddressesForAccount", reflect.TypeOf((*MockStore)(nil).GetActiveProxyClusterAddressesForAccount), ctx, accountID) } -// GetActiveProxyClusters mocks base method. -func (m *MockStore) GetActiveProxyClusters(ctx context.Context, accountID string) ([]proxy.Cluster, error) { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "GetActiveProxyClusters", ctx, accountID) - ret0, _ := ret[0].([]proxy.Cluster) - ret1, _ := ret[1].(error) - return ret0, ret1 -} - -// GetActiveProxyClusters indicates an expected call of GetActiveProxyClusters. -func (mr *MockStoreMockRecorder) GetActiveProxyClusters(ctx, accountID interface{}) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveProxyClusters", reflect.TypeOf((*MockStore)(nil).GetActiveProxyClusters), ctx, accountID) -} - // GetAllAccounts mocks base method. func (m *MockStore) GetAllAccounts(ctx context.Context) []*types2.Account { m.ctrl.T.Helper() @@ -2048,6 +2047,21 @@ func (mr *MockStoreMockRecorder) GetProxyByAccountID(ctx, accountID interface{}) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetProxyByAccountID", reflect.TypeOf((*MockStore)(nil).GetProxyByAccountID), ctx, accountID) } +// GetProxyClusters mocks base method. +func (m *MockStore) GetProxyClusters(ctx context.Context, accountID string) ([]proxy.Cluster, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetProxyClusters", ctx, accountID) + ret0, _ := ret[0].([]proxy.Cluster) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetProxyClusters indicates an expected call of GetProxyClusters. +func (mr *MockStoreMockRecorder) GetProxyClusters(ctx, accountID interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetProxyClusters", reflect.TypeOf((*MockStore)(nil).GetProxyClusters), ctx, accountID) +} + // GetResourceGroups mocks base method. func (m *MockStore) GetResourceGroups(ctx context.Context, lockStrength LockingStrength, accountID, resourceID string) ([]*types2.Group, error) { m.ctrl.T.Helper() @@ -2950,20 +2964,6 @@ func (mr *MockStoreMockRecorder) SaveProxy(ctx, proxy interface{}) *gomock.Call return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SaveProxy", reflect.TypeOf((*MockStore)(nil).SaveProxy), ctx, proxy) } -// DisconnectProxy mocks base method. -func (m *MockStore) DisconnectProxy(ctx context.Context, proxyID, sessionID string) error { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "DisconnectProxy", ctx, proxyID, sessionID) - ret0, _ := ret[0].(error) - return ret0 -} - -// DisconnectProxy indicates an expected call of DisconnectProxy. -func (mr *MockStoreMockRecorder) DisconnectProxy(ctx, proxyID, sessionID interface{}) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DisconnectProxy", reflect.TypeOf((*MockStore)(nil).DisconnectProxy), ctx, proxyID, sessionID) -} - // SaveProxyAccessToken mocks base method. func (m *MockStore) SaveProxyAccessToken(ctx context.Context, token *types2.ProxyAccessToken) error { m.ctrl.T.Helper() diff --git a/proxy/management_integration_test.go b/proxy/management_integration_test.go index 9fd3d2ce9..d7e891801 100644 --- a/proxy/management_integration_test.go +++ b/proxy/management_integration_test.go @@ -366,7 +366,7 @@ func (m *storeBackedServiceManager) GetServiceByDomain(ctx context.Context, doma return m.store.GetServiceByDomain(ctx, domain) } -func (m *storeBackedServiceManager) GetActiveClusters(_ context.Context, _, _ string) ([]nbproxy.Cluster, error) { +func (m *storeBackedServiceManager) GetClusters(_ context.Context, _, _ string) ([]nbproxy.Cluster, error) { return nil, nil } diff --git a/shared/management/http/api/openapi.yml b/shared/management/http/api/openapi.yml index 942f3aa45..353aff72d 100644 --- a/shared/management/http/api/openapi.yml +++ b/shared/management/http/api/openapi.yml @@ -3417,19 +3417,43 @@ components: type: string description: Cluster address used for CNAME targets example: "eu.proxy.netbird.io" + type: + $ref: '#/components/schemas/ProxyClusterType' + online: + type: boolean + description: Whether at least one proxy in the cluster has heartbeated within the active window + example: true connected_proxies: type: integer - description: Number of proxy nodes connected in this cluster + description: Number of proxy nodes currently connected (heartbeat within the active window) example: 3 - self_hosted: + supports_custom_ports: type: boolean - description: Whether this cluster is a self-hosted (BYOP) proxy managed by the account owner + description: Whether the cluster supports binding arbitrary TCP/UDP ports + example: true + require_subdomain: + type: boolean + description: Whether services on this cluster must include a subdomain label + example: false + supports_crowdsec: + type: boolean + description: Whether all active proxies in the cluster have CrowdSec configured example: false required: - id - address + - type + - online - connected_proxies - - self_hosted + ProxyClusterType: + type: string + description: | + Source of the proxy cluster. `account` clusters are owned and operated by the account (BYOP); + `shared` clusters are operated by NetBird and shared across accounts. + enum: + - account + - shared + example: shared ReverseProxyDomainType: type: string description: Type of Reverse Proxy Domain diff --git a/shared/management/http/api/types.gen.go b/shared/management/http/api/types.gen.go index b3bb475a9..16e765f8c 100644 --- a/shared/management/http/api/types.gen.go +++ b/shared/management/http/api/types.gen.go @@ -1,6 +1,6 @@ // Package api provides primitives to interact with the openapi HTTP API. // -// Code generated by github.com/oapi-codegen/oapi-codegen/v2 version v2.6.0 DO NOT EDIT. +// Code generated by github.com/oapi-codegen/oapi-codegen/v2 version v2.7.0 DO NOT EDIT. package api import ( @@ -13,8 +13,8 @@ import ( ) const ( - BearerAuthScopes = "BearerAuth.Scopes" - TokenAuthScopes = "TokenAuth.Scopes" + BearerAuthScopes bearerAuthContextKey = "BearerAuth.Scopes" + TokenAuthScopes tokenAuthContextKey = "TokenAuth.Scopes" ) // Defines values for AccessRestrictionsCrowdsecMode. @@ -511,6 +511,7 @@ func (e GroupMinimumIssued) Valid() bool { // Defines values for IdentityProviderType. const ( + IdentityProviderTypeAdfs IdentityProviderType = "adfs" IdentityProviderTypeEntra IdentityProviderType = "entra" IdentityProviderTypeGoogle IdentityProviderType = "google" IdentityProviderTypeMicrosoft IdentityProviderType = "microsoft" @@ -518,12 +519,13 @@ const ( IdentityProviderTypeOkta IdentityProviderType = "okta" IdentityProviderTypePocketid IdentityProviderType = "pocketid" IdentityProviderTypeZitadel IdentityProviderType = "zitadel" - IdentityProviderTypeAdfs IdentityProviderType = "adfs" ) // Valid indicates whether the value is a known member of the IdentityProviderType enum. func (e IdentityProviderType) Valid() bool { switch e { + case IdentityProviderTypeAdfs: + return true case IdentityProviderTypeEntra: return true case IdentityProviderTypeGoogle: @@ -538,8 +540,6 @@ func (e IdentityProviderType) Valid() bool { return true case IdentityProviderTypeZitadel: return true - case IdentityProviderTypeAdfs: - return true default: return false } @@ -878,6 +878,24 @@ func (e PolicyRuleUpdateProtocol) Valid() bool { } } +// Defines values for ProxyClusterType. +const ( + ProxyClusterTypeAccount ProxyClusterType = "account" + ProxyClusterTypeShared ProxyClusterType = "shared" +) + +// Valid indicates whether the value is a known member of the ProxyClusterType enum. +func (e ProxyClusterType) Valid() bool { + switch e { + case ProxyClusterTypeAccount: + return true + case ProxyClusterTypeShared: + return true + default: + return false + } +} + // Defines values for ResourceType. const ( ResourceTypeDomain ResourceType = "domain" @@ -1638,7 +1656,9 @@ type Checks struct { // OsVersionCheck Posture check for the version of operating system OsVersionCheck *OSVersionCheck `json:"os_version_check,omitempty"` - // PeerNetworkRangeCheck Posture check for allow or deny access based on the peer's IP addresses. A range matches when it contains any of the peer's local network interface IPs or its public connection (NAT egress) IP, so ranges may target private subnets, public CIDRs, or single hosts via a /32 or /128. + // PeerNetworkRangeCheck Posture check for allow or deny access based on the peer's IP addresses. A range matches when it + // contains any of the peer's local network interface IPs or its public connection (NAT egress) IP, + // so ranges may target private subnets, public CIDRs, or single hosts via a /32 or /128. PeerNetworkRangeCheck *PeerNetworkRangeCheck `json:"peer_network_range_check,omitempty"` // ProcessCheck Posture Check for binaries exist and are running in the peer’s system @@ -3330,7 +3350,9 @@ type PeerMinimum struct { Name string `json:"name"` } -// PeerNetworkRangeCheck Posture check for allow or deny access based on the peer's IP addresses. A range matches when it contains any of the peer's local network interface IPs or its public connection (NAT egress) IP, so ranges may target private subnets, public CIDRs, or single hosts via a /32 or /128. +// PeerNetworkRangeCheck Posture check for allow or deny access based on the peer's IP addresses. A range matches when it +// contains any of the peer's local network interface IPs or its public connection (NAT egress) IP, +// so ranges may target private subnets, public CIDRs, or single hosts via a /32 or /128. type PeerNetworkRangeCheck struct { // Action Action to take upon policy match Action PeerNetworkRangeCheckAction `json:"action"` @@ -3785,19 +3807,36 @@ type ProxyAccessLogsResponse struct { // ProxyCluster A proxy cluster represents a group of proxy nodes serving the same address type ProxyCluster struct { - // Id Unique identifier of a proxy in this cluster - Id string `json:"id"` - // Address Cluster address used for CNAME targets Address string `json:"address"` - // ConnectedProxies Number of proxy nodes connected in this cluster + // ConnectedProxies Number of proxy nodes currently connected (heartbeat within the active window) ConnectedProxies int `json:"connected_proxies"` - // SelfHosted Whether this cluster is a self-hosted (BYOP) proxy managed by the account owner - SelfHosted bool `json:"self_hosted"` + // Id Unique identifier of a proxy in this cluster + Id string `json:"id"` + + // Online Whether at least one proxy in the cluster has heartbeated within the active window + Online bool `json:"online"` + + // RequireSubdomain Whether services on this cluster must include a subdomain label + RequireSubdomain *bool `json:"require_subdomain,omitempty"` + + // SupportsCrowdsec Whether all active proxies in the cluster have CrowdSec configured + SupportsCrowdsec *bool `json:"supports_crowdsec,omitempty"` + + // SupportsCustomPorts Whether the cluster supports binding arbitrary TCP/UDP ports + SupportsCustomPorts *bool `json:"supports_custom_ports,omitempty"` + + // Type Source of the proxy cluster. `account` clusters are owned and operated by the account (BYOP); + // `shared` clusters are operated by NetBird and shared across accounts. + Type ProxyClusterType `json:"type"` } +// ProxyClusterType Source of the proxy cluster. `account` clusters are owned and operated by the account (BYOP); +// `shared` clusters are operated by NetBird and shared across accounts. +type ProxyClusterType string + // ProxyToken defines model for ProxyToken. type ProxyToken struct { CreatedAt time.Time `json:"created_at"` @@ -4820,6 +4859,12 @@ type ZoneRequest struct { // Conflict Standard error response. Note: The exact structure of this error response is inferred from `util.WriteErrorResponse` and `util.WriteError` usage in the provided Go code, as a specific Go struct for errors was not provided. type Conflict = ErrorResponse +// bearerAuthContextKey is the context key for BearerAuth security scheme +type bearerAuthContextKey string + +// tokenAuthContextKey is the context key for TokenAuth security scheme +type tokenAuthContextKey string + // GetApiEventsNetworkTrafficParams defines parameters for GetApiEventsNetworkTraffic. type GetApiEventsNetworkTrafficParams struct { // Page Page number From c784b0255063b9cbfde830c78670de2400e46c1c Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Wed, 20 May 2026 12:21:03 +0200 Subject: [PATCH 26/28] [misc] Update contribution guidelines (#6219) Update contribution guidelines and PR template to require discussing impactful changes with the team --- .github/pull_request_template.md | 1 + CONTRIBUTING.md | 9 +++++++++ 2 files changed, 10 insertions(+) diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 9d6bc96eb..8e68054bd 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -12,6 +12,7 @@ - [ ] Is a feature enhancement - [ ] It is a refactor - [ ] Created tests that fail without the change (if possible) +- [ ] This change does **not** modify the public API, gRPC protocols, functionality behavior, CLI / service flags, or introduce a new feature — **OR** I have discussed it with the NetBird team beforehand (link the issue / Slack thread in the description). See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#discuss-changes-with-the-netbird-team-first). > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 960cd30e9..cd1c087bb 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -15,6 +15,7 @@ If you haven't already, join our slack workspace [here](https://docs.netbird.io/ - [Contributing to NetBird](#contributing-to-netbird) - [Contents](#contents) - [Code of conduct](#code-of-conduct) + - [Discuss changes with the NetBird team first](#discuss-changes-with-the-netbird-team-first) - [Directory structure](#directory-structure) - [Development setup](#development-setup) - [Requirements](#requirements) @@ -33,6 +34,14 @@ Conduct which can be found in the file [CODE_OF_CONDUCT.md](CODE_OF_CONDUCT.md). By participating, you are expected to uphold this code. Please report unacceptable behavior to community@netbird.io. +## Discuss changes with the NetBird team first + +Changes to the **public API**, **gRPC protocols**, **functionality behavior**, **CLI / service flags**, or **new features** should be discussed with the NetBird team before you start the work. These surfaces are part of NetBird's contract with operators, self-hosters, and downstream integrators, and changes to them have compatibility, security, and release-planning implications that benefit from an early conversation. + +Open an issue or reach out on [Slack](https://docs.netbird.io/slack-url) to talk through what you have in mind. We'll help shape the change, flag any constraints we know about, and confirm the direction so the PR review can focus on implementation rather than design. + +Typical bug fixes, internal refactors, documentation updates, and tests do not need pre-discussion — open the PR directly. + ## Directory structure The NetBird project monorepo is organized to maintain most of its individual dependencies code within their directories, except for a few auxiliary or shared packages. From 9192b4f029f8f0eeaad77fff8625c34ba9849668 Mon Sep 17 00:00:00 2001 From: Viktor Liu <17948409+lixmal@users.noreply.github.com> Date: Wed, 20 May 2026 20:09:22 +0900 Subject: [PATCH 27/28] [client] Bump macOS sleep callback timeout to 20s (#6220) --- client/internal/sleep/detector_darwin.go | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/client/internal/sleep/detector_darwin.go b/client/internal/sleep/detector_darwin.go index ef495bded..fc4713b21 100644 --- a/client/internal/sleep/detector_darwin.go +++ b/client/internal/sleep/detector_darwin.go @@ -188,7 +188,9 @@ func (d *Detector) triggerCallback(event EventType, cb func(event EventType), do } doneChan := make(chan struct{}) - timeout := time.NewTimer(500 * time.Millisecond) + // macOS forces sleep ~30s after kIOMessageSystemWillSleep, so block long + // enough for teardown to finish while staying under that deadline. + timeout := time.NewTimer(20 * time.Second) defer timeout.Stop() go func() { From 4955c345d53f63394266305744841c4e1bff8123 Mon Sep 17 00:00:00 2001 From: Viktor Liu <17948409+lixmal@users.noreply.github.com> Date: Wed, 20 May 2026 23:25:56 +0900 Subject: [PATCH 28/28] Clean up README header, key features table, and self-hosted quickstart (#6178) --- README.md | 153 +++++++++++++++++++++++++----------------------------- 1 file changed, 70 insertions(+), 83 deletions(-) diff --git a/README.md b/README.md index dc84af2fd..cc27e2d28 100644 --- a/README.md +++ b/README.md @@ -1,147 +1,134 @@
-
-
-

- -

-

- - - - - - -
+

+ NetBird logo +

+

+ + SonarCloud alert status + + + BSD-3 License + - - + NetBird Slack + - - -
+ Community forum + - - + Gurubase: Ask NetBird Guru +

-

- - Start using NetBird at netbird.io + + Start using NetBird at netbird.io +
+ See Documentation +
+ Join our Slack channel or our Community forum +

- See Documentation
- Join our Slack channel or our Community forum -
- -
-
- - 🚀 We are hiring! Join us at careers.netbird.io - -
-
- - New: NetBird terraform provider - + + 🚀 We are hiring! Join us at careers.netbird.io +

-
- **NetBird combines a configuration-free peer-to-peer private network and a centralized access control system in a single platform, making it easy to create secure private networks for your organization or home.** **Connect.** NetBird creates a WireGuard-based overlay network that automatically connects your machines over an encrypted tunnel, leaving behind the hassle of opening ports, complex firewall rules, VPN gateways, and so forth. **Secure.** NetBird enables secure remote access by applying granular access policies while allowing you to manage them intuitively from a single place. Works universally on any infrastructure. -### Open Source Network Security in a Single Platform - https://github.com/user-attachments/assets/10cec749-bb56-4ab3-97af-4e38850108d2 -### Self-Host NetBird (Video) +### Self-host NetBird (video) + [![Watch the video](https://img.youtube.com/vi/bZAgpT6nzaQ/0.jpg)](https://youtu.be/bZAgpT6nzaQ) ### Key features -| Connectivity | Management | Security | Automation| Platforms | -|----|----|----|----|----| -|
  • - \[x] Kernel WireGuard
|
  • - \[x] [Admin Web UI](https://github.com/netbirdio/dashboard)
|
  • - \[x] [SSO & MFA support](https://docs.netbird.io/how-to/installation#running-net-bird-with-sso-login)
|
  • - \[x] [Public API](https://docs.netbird.io/api)
|
  • - \[x] Linux
| -|
  • - \[x] Peer-to-peer connections
|
  • - \[x] Auto peer discovery and configuration
  • |
    • - \[x] [Access control - groups & rules](https://docs.netbird.io/how-to/manage-network-access)
    • |
      • - \[x] [Setup keys for bulk network provisioning](https://docs.netbird.io/how-to/register-machines-using-setup-keys)
      • |
        • - \[x] Mac
        • | -|
          • - \[x] Connection relay fallback
          • |
            • - \[x] [IdP integrations](https://docs.netbird.io/selfhosted/identity-providers)
            • |
              • - \[x] [Activity logging](https://docs.netbird.io/how-to/audit-events-logging)
              • |
                • - \[x] [Self-hosting quickstart script](https://docs.netbird.io/selfhosted/selfhosted-quickstart)
                • |
                  • - \[x] Windows
                  • | -|
                    • - \[x] [Routes to external networks](https://docs.netbird.io/how-to/routing-traffic-to-private-networks)
                    • |
                      • - \[x] [Private DNS](https://docs.netbird.io/how-to/manage-dns-in-your-network)
                      • |
                        • - \[x] [Device posture checks](https://docs.netbird.io/how-to/manage-posture-checks)
                        • |
                          • - \[x] IdP groups sync with JWT
                          • |
                            • - \[x] Android
                            • | -|
                              • - \[x] NAT traversal with BPF
                              • |
                                • - \[x] [Multiuser support](https://docs.netbird.io/how-to/add-users-to-your-network)
                                • |
                                  • - \[x] Peer-to-peer encryption
                                  • ||
                                    • - \[x] iOS
                                    • | -|||
                                      • - \[x] [Quantum-resistance with Rosenpass](https://netbird.io/knowledge-hub/the-first-quantum-resistant-mesh-vpn)
                                      • ||
                                        • - \[x] OpenWRT
                                        • | -|||
                                          • - \[x] [Periodic re-authentication](https://docs.netbird.io/how-to/enforce-periodic-user-authentication)
                                          • ||
                                            • - \[x] [Serverless](https://docs.netbird.io/how-to/netbird-on-faas)
                                            • | -|||||
                                              • - \[x] Docker
                                              • | +| Connectivity | Management | Security | Automation | Platforms | +|---|---|---|---|---| +| ✓ [Kernel WireGuard](https://docs.netbird.io/about-netbird/why-wireguard-with-netbird) | ✓ [Admin Web UI](https://github.com/netbirdio/dashboard) | ✓ [SSO & MFA support](https://docs.netbird.io/how-to/installation#running-net-bird-with-sso-login) | ✓ [Public API](https://docs.netbird.io/api) | ✓ [Linux](https://docs.netbird.io/get-started/install/linux) | +| ✓ [Peer-to-peer connections](https://docs.netbird.io/about-netbird/how-netbird-works) | ✓ Auto peer discovery and configuration | ✓ [Access control: groups & rules](https://docs.netbird.io/how-to/manage-network-access) | ✓ [Setup keys for bulk provisioning](https://docs.netbird.io/how-to/register-machines-using-setup-keys) | ✓ [macOS](https://docs.netbird.io/get-started/install/macos) | +| ✓ Connection relay fallback | ✓ [IdP integrations](https://docs.netbird.io/selfhosted/identity-providers) | ✓ [Activity logging](https://docs.netbird.io/how-to/audit-events-logging) | ✓ [Self-hosting quickstart script](https://docs.netbird.io/selfhosted/selfhosted-quickstart) | ✓ [Windows](https://docs.netbird.io/get-started/install/windows) | +| ✓ [Routes to external networks](https://docs.netbird.io/how-to/routing-traffic-to-private-networks) | ✓ [Private DNS](https://docs.netbird.io/how-to/manage-dns-in-your-network) | ✓ [Traffic events](https://docs.netbird.io/manage/activity/traffic-events-logging) | ✓ [IdP groups sync with JWT](https://docs.netbird.io/manage/team/idp-sync) | ✓ [Android](https://docs.netbird.io/get-started/install/android) | +| ✓ [Domain-based DNS routes](https://docs.netbird.io/manage/dns/dns-aliases-for-routed-networks) | ✓ [Custom DNS zones](https://docs.netbird.io/manage/dns/custom-zones) | ✓ [Device posture checks](https://docs.netbird.io/how-to/manage-posture-checks) | ✓ [Terraform provider](https://registry.terraform.io/providers/netbirdio/netbird/latest) | ✓ [Android TV](https://docs.netbird.io/get-started/install/android-tv) | +| ✓ [Exit nodes](https://docs.netbird.io/manage/network-routes/use-cases/exit-nodes) | ✓ [Multiuser support](https://docs.netbird.io/how-to/add-users-to-your-network) | ✓ Peer-to-peer encryption | ✓ [Ansible collection](https://github.com/netbirdio/ansible-netbird) | ✓ [iOS](https://docs.netbird.io/get-started/install/ios) | +| ✓ [IPv6 dual-stack overlay](https://docs.netbird.io/manage/settings/ipv6) | ✓ [Multi-account profile switching](https://docs.netbird.io/client/profiles) | ✓ [SSH with central access policies](https://docs.netbird.io/manage/peers/ssh) | | ✓ [Apple TV](https://docs.netbird.io/get-started/install/tvos) | +| ✓ [Browser SSH & RDP](https://docs.netbird.io/manage/peers/browser-client) | | ✓ [Quantum-resistance with Rosenpass](https://netbird.io/knowledge-hub/the-first-quantum-resistant-mesh-vpn) | | ✓ FreeBSD | +| ✓ [Reverse proxy with auto-TLS](https://docs.netbird.io/manage/reverse-proxy) | | ✓ [Periodic re-authentication](https://docs.netbird.io/how-to/enforce-periodic-user-authentication) | | ✓ [pfSense](https://docs.netbird.io/get-started/install/pfsense) | +| | | | | ✓ [OPNsense](https://docs.netbird.io/get-started/install/opnsense) | +| | | | | ✓ [MikroTik RouterOS](https://docs.netbird.io/use-cases/homelab/client-on-mikrotik-router) | +| | | | | ✓ OpenWRT | +| | | | | ✓ [Synology](https://docs.netbird.io/get-started/install/synology) | +| | | | | ✓ [TrueNAS](https://docs.netbird.io/get-started/install/truenas) | +| | | | | ✓ [Proxmox](https://docs.netbird.io/get-started/install/proxmox-ve) | +| | | | | ✓ [Raspberry Pi](https://docs.netbird.io/get-started/install/raspberrypi) | +| | | | | ✓ [Serverless](https://docs.netbird.io/how-to/netbird-on-faas) | +| | | | | ✓ [Container](https://docs.netbird.io/get-started/install/docker) | ### Quickstart with NetBird Cloud -- Download and install NetBird at [https://app.netbird.io/install](https://app.netbird.io/install) -- Follow the steps to sign-up with Google, Microsoft, GitHub or your email address. -- Check NetBird [admin UI](https://app.netbird.io/). -- Add more machines. +- Download and install NetBird at [https://app.netbird.io/install](https://app.netbird.io/install). +- Follow the steps to sign up with Google, Microsoft, GitHub or your email address. +- Check the NetBird [admin UI](https://app.netbird.io/). ### Quickstart with self-hosted NetBird -> This is the quickest way to try self-hosted NetBird. It should take around 5 minutes to get started if you already have a public domain and a VM. -Follow the [Advanced guide with a custom identity provider](https://docs.netbird.io/selfhosted/selfhosted-guide#advanced-guide-with-a-custom-identity-provider) for installations with different IDPs. +This is the quickest way to try self-hosted NetBird. It should take around 5 minutes to get started if you already have a public domain and a VM. Follow the [Advanced guide with a custom identity provider](https://docs.netbird.io/selfhosted/selfhosted-guide#advanced-guide-with-a-custom-identity-provider) for installations with different IdPs. **Infrastructure requirements:** -- A Linux VM with at least **1CPU** and **2GB** of memory. -- The VM should be publicly accessible on TCP ports **80** and **443** and UDP port: **3478**. -- **Public domain** name pointing to the VM. +- A Linux VM with at least **1 CPU** and **2 GB** of memory. +- The VM should be publicly accessible on TCP ports **80** and **443** and UDP port **3478**. +- A **public domain** name pointing to the VM. **Software requirements:** -- Docker installed on the VM with the docker-compose plugin ([Docker installation guide](https://docs.docker.com/engine/install/)) or docker with docker-compose in version 2 or higher. -- [jq](https://jqlang.github.io/jq/) installed. In most distributions - Usually available in the official repositories and can be installed with `sudo apt install jq` or `sudo yum install jq` -- [curl](https://curl.se/) installed. - Usually available in the official repositories and can be installed with `sudo apt install curl` or `sudo yum install curl` +- Docker with the Compose plugin (Compose v2 or higher). See the [Docker installation guide](https://docs.docker.com/engine/install/). **Steps** - Download and run the installation script: ```bash export NETBIRD_DOMAIN=netbird.example.com; curl -fsSL https://github.com/netbirdio/netbird/releases/latest/download/getting-started.sh | bash ``` -- Once finished, you can manage the resources via `docker-compose` ### A bit on NetBird internals -- Every machine in the network runs [NetBird Agent (or Client)](client/) that manages WireGuard. -- Every agent connects to [Management Service](management/) that holds network state, manages peer IPs, and distributes network updates to agents (peers). -- NetBird agent uses WebRTC ICE implemented in [pion/ice library](https://github.com/pion/ice) to discover connection candidates when establishing a peer-to-peer connection between machines. -- Connection candidates are discovered with the help of [STUN](https://en.wikipedia.org/wiki/STUN) servers. -- Agents negotiate a connection through [Signal Service](signal/) passing p2p encrypted messages with candidates. -- Sometimes the NAT traversal is unsuccessful due to strict NATs (e.g. mobile carrier-grade NAT) and a p2p connection isn't possible. When this occurs the system falls back to a relay server called [TURN](https://en.wikipedia.org/wiki/Traversal_Using_Relays_around_NAT), and a secure WireGuard tunnel is established via the TURN server. - -[Coturn](https://github.com/coturn/coturn) is the one that has been successfully used for STUN and TURN in NetBird setups. +- Every machine in the network runs the [NetBird agent](client/), which manages WireGuard. +- Every agent connects to the [Management Service](management/), which holds network state, manages peer IPs, and distributes updates to agents. +- Agents use ICE (via [pion/ice](https://github.com/pion/ice)) to discover connection candidates for peer-to-peer connections. +- Candidates are discovered with the help of [STUN](https://en.wikipedia.org/wiki/STUN) servers. +- Agents negotiate a connection through the [Signal Service](signal/), exchanging end-to-end encrypted messages with candidates. +- When NAT traversal fails (e.g. mobile carrier-grade NAT) and a direct p2p connection isn't possible, the system falls back to a [Relay Service](relay/) and a secure WireGuard tunnel is established through it.

                                                - + NetBird high-level architecture diagram

                                                See a complete [architecture overview](https://docs.netbird.io/about-netbird/how-netbird-works#architecture) for details. ### Community projects -- [NetBird installer script](https://github.com/physk/netbird-installer) -- [NetBird ansible collection by Dominion Solutions](https://galaxy.ansible.com/ui/repo/published/dominion_solutions/netbird/) -- [netbird-tui](https://github.com/n0pashkov/netbird-tui) — terminal UI for managing NetBird peers, routes, and settings +- [NetBird installer script](https://github.com/physk/netbird-installer) +- [netbird-tui](https://github.com/n0pashkov/netbird-tui) - terminal UI for managing NetBird peers, routes, and settings +- [caddy-netbird](https://github.com/lixmal/caddy-netbird) - Caddy plugin that embeds a NetBird client for proxying HTTP and TCP/UDP traffic through NetBird networks **Note**: The `main` branch may be in an *unstable or even broken state* during development. For stable versions, see [releases](https://github.com/netbirdio/netbird/releases). ### Support acknowledgement -In November 2022, NetBird joined the [StartUpSecure program](https://www.forschung-it-sicherheit-kommunikationssysteme.de/foerderung/bekanntmachungen/startup-secure) sponsored by The Federal Ministry of Education and Research of The Federal Republic of Germany. Together with [CISPA Helmholtz Center for Information Security](https://cispa.de/en) NetBird brings the security best practices and simplicity to private networking. +In November 2022, NetBird joined the [StartUpSecure program](https://www.forschung-it-sicherheit-kommunikationssysteme.de/foerderung/bekanntmachungen/startup-secure) sponsored by the Federal Ministry of Education and Research of the Federal Republic of Germany. Together with the [CISPA Helmholtz Center for Information Security](https://cispa.de/en), NetBird brings security best practices and simplicity to private networking. ![CISPA_Logo_BLACK_EN_RZ_RGB (1)](https://user-images.githubusercontent.com/700848/203091324-c6d311a0-22b5-4b05-a288-91cbc6cdcc46.png) -### Testimonials -We use open-source technologies like [WireGuard®](https://www.wireguard.com/), [Pion ICE (WebRTC)](https://github.com/pion/ice), [Coturn](https://github.com/coturn/coturn), and [Rosenpass](https://rosenpass.eu). We very much appreciate the work these guys are doing and we'd greatly appreciate if you could support them in any way (e.g., by giving a star or a contribution). +### Acknowledgements +We build on open-source technologies like [WireGuard®](https://www.wireguard.com/), [Pion ICE](https://github.com/pion/ice), and [Rosenpass](https://rosenpass.eu). We greatly appreciate the work these projects are doing, and we'd love it if you could support them too (e.g., by starring or contributing). ### Legal -This repository is licensed under BSD-3-Clause license that applies to all parts of the repository except for the directories management/, signal/ and relay/. +This repository is licensed under the BSD-3-Clause license, which applies to all parts of the repository except for the directories management/, signal/ and relay/. Those directories are licensed under the GNU Affero General Public License version 3.0 (AGPLv3). See the respective LICENSE files inside each directory. _WireGuard_ and the _WireGuard_ logo are [registered trademarks](https://www.wireguard.com/trademark-policy/) of Jason A. Donenfeld.