From c2471510b675e6d294a645b858b0947dac14d0cc Mon Sep 17 00:00:00 2001 From: pascal Date: Wed, 3 Jun 2026 18:31:38 +0200 Subject: [PATCH] use unified resolver --- management/server/account/manager.go | 4 +- management/server/account/manager_mock.go | 15 + management/server/affected_groups.go | 425 ----------- .../server/affected_peers_coverage_test.go | 119 ++++ .../server/affected_peers_oldstate_test.go | 143 ++++ .../server/affected_peers_property_test.go | 251 +++++++ .../server/affected_peers_querycount_test.go | 125 ++++ .../affected_peers_router_paths_test.go | 369 ++++++++++ .../server/affected_peers_router_test.go | 7 +- management/server/affected_peers_test.go | 230 +----- management/server/affectedpeers/resolver.go | 661 ++++++++++++++++++ .../server/affectedpeers/resolver_test.go | 138 ++++ management/server/group.go | 35 +- management/server/group_linkage.go | 9 + management/server/mock_server/account_mock.go | 9 + management/server/networks/manager.go | 152 +--- .../server/networks/resources/manager.go | 195 +----- management/server/networks/routers/manager.go | 224 ++---- management/server/peer.go | 25 +- management/server/policy.go | 42 +- management/server/posture_checks.go | 25 +- management/server/route.go | 29 +- 22 files changed, 1976 insertions(+), 1256 deletions(-) delete mode 100644 management/server/affected_groups.go create mode 100644 management/server/affected_peers_coverage_test.go create mode 100644 management/server/affected_peers_oldstate_test.go create mode 100644 management/server/affected_peers_property_test.go create mode 100644 management/server/affected_peers_querycount_test.go create mode 100644 management/server/affected_peers_router_paths_test.go create mode 100644 management/server/affectedpeers/resolver.go create mode 100644 management/server/affectedpeers/resolver_test.go diff --git a/management/server/account/manager.go b/management/server/account/manager.go index bc5afbe81..2e6a11f05 100644 --- a/management/server/account/manager.go +++ b/management/server/account/manager.go @@ -13,6 +13,7 @@ import ( nbdns "github.com/netbirdio/netbird/dns" "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/affectedpeers" nbcache "github.com/netbirdio/netbird/management/server/cache" "github.com/netbirdio/netbird/management/server/idp" nbpeer "github.com/netbirdio/netbird/management/server/peer" @@ -109,7 +110,7 @@ type Manager interface { UpdateAccountSettings(ctx context.Context, accountID, userID string, newSettings *types.Settings) (*types.Settings, error) UpdateAccountOnboarding(ctx context.Context, accountID, userID string, newOnboarding *types.AccountOnboarding) (*types.AccountOnboarding, error) LoginPeer(ctx context.Context, login types.PeerLogin) (*nbpeer.Peer, *types.NetworkMap, []*posture.Checks, error) // used by peer gRPC API - ExtendPeerSession(ctx context.Context, peerPubKey, userID string) (time.Time, error) // used by peer gRPC API for ExtendAuthSession + ExtendPeerSession(ctx context.Context, peerPubKey, userID string) (time.Time, error) // used by peer gRPC API for ExtendAuthSession SyncPeer(ctx context.Context, sync types.PeerSync, accountID string) (*nbpeer.Peer, *types.NetworkMap, []*posture.Checks, int64, error) // used by peer gRPC API GetExternalCacheManager() ExternalCacheManager GetPostureChecks(ctx context.Context, accountID, postureChecksID, userID string) (*posture.Checks, error) @@ -130,6 +131,7 @@ type Manager interface { 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, reason types.UpdateReason) + ResolveAffectedPeers(ctx context.Context, s store.Store, accountID string, change affectedpeers.Change) []string 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 3f7fd2004..dcd9a671f 100644 --- a/management/server/account/manager_mock.go +++ b/management/server/account/manager_mock.go @@ -15,6 +15,7 @@ import ( dns "github.com/netbirdio/netbird/dns" service "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" activity "github.com/netbirdio/netbird/management/server/activity" + affectedpeers "github.com/netbirdio/netbird/management/server/affectedpeers" idp "github.com/netbirdio/netbird/management/server/idp" peer "github.com/netbirdio/netbird/management/server/peer" posture "github.com/netbirdio/netbird/management/server/posture" @@ -1661,6 +1662,20 @@ func (mr *MockManagerMockRecorder) UpdateAffectedPeers(ctx, accountID, peerIDs i return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateAffectedPeers", reflect.TypeOf((*MockManager)(nil).UpdateAffectedPeers), ctx, accountID, peerIDs) } +// ResolveAffectedPeers mocks base method. +func (m *MockManager) ResolveAffectedPeers(ctx context.Context, s store.Store, accountID string, change affectedpeers.Change) []string { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ResolveAffectedPeers", ctx, s, accountID, change) + ret0, _ := ret[0].([]string) + return ret0 +} + +// ResolveAffectedPeers indicates an expected call of ResolveAffectedPeers. +func (mr *MockManagerMockRecorder) ResolveAffectedPeers(ctx, s, accountID, change interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ResolveAffectedPeers", reflect.TypeOf((*MockManager)(nil).ResolveAffectedPeers), ctx, s, accountID, change) +} + // 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/affected_groups.go b/management/server/affected_groups.go deleted file mode 100644 index c5f108ea9..000000000 --- a/management/server/affected_groups.go +++ /dev/null @@ -1,425 +0,0 @@ -package server - -import ( - "context" - - log "github.com/sirupsen/logrus" - - rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" - 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" -) - -// 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 - } - - changedGroupSet := toSet(changedGroupIDs) - changedPeerSet := toSet(changedPeerIDs) - - groupSet := make(map[string]struct{}) - peerSet := make(map[string]struct{}) - - 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) - collectAffectedFromProxyServices(ctx, transaction, accountID, changedGroupSet, changedPeerSet, peerSet) - - allGroupIDs = setToSlice(groupSet) - directPeerIDs = setToSlice(peerSet) - - log.WithContext(ctx).Tracef("affected groups resolution: changedGroups=%v changedPeers=%v -> affectedGroups=%v, directPeers=%v", - changedGroupIDs, changedPeerIDs, allGroupIDs, directPeerIDs) - - return allGroupIDs, directPeerIDs -} - -// 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 affected group resolution: %v", err) - return - } - - for _, policy := range policies { - matchedByGroup := policyReferencesGroups(policy, changedGroupSet) - matchedByPeer := len(changedPeerSet) > 0 && policyReferencesDirectPeers(policy, changedPeerSet) - if !matchedByGroup && !matchedByPeer { - continue - } - 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{}{} - } - } -} - -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{}{} - } - } -} - -// collectAffectedFromProxyServices handles policies that are synthesized at -// network-map computation time by Account.InjectProxyPolicies. Those policies -// connect proxy peers (peer.ProxyMeta.Embedded == true) to service targets and -// never reach the database, so the other collectors cannot see them. -func collectAffectedFromProxyServices(ctx context.Context, transaction store.Store, accountID string, changedGroupSet, changedPeerSet map[string]struct{}, peerSet map[string]struct{}) { - if len(changedGroupSet) == 0 && len(changedPeerSet) == 0 { - return - } - - services, proxyByCluster, ok := loadProxyServiceContext(ctx, transaction, accountID) - if !ok { - return - } - - expandedPeerSet := expandChangedPeersWithGroups(ctx, transaction, accountID, changedGroupSet, changedPeerSet) - - for _, svc := range services { - if svc == nil { - continue - } - proxyPeers := proxyByCluster[svc.ProxyCluster] - if len(proxyPeers) == 0 { - continue - } - if !serviceMatchesChangedPeers(svc, proxyPeers, expandedPeerSet) { - continue - } - - log.WithContext(ctx).Tracef("collectAffectedFromProxyServices: service %s (cluster=%s) matched; folding %d proxy peers and target peers", - svc.ID, svc.ProxyCluster, len(proxyPeers)) - addServicePeersToSet(svc, proxyPeers, peerSet) - } -} - -func loadProxyServiceContext(ctx context.Context, transaction store.Store, accountID string) ([]*rpservice.Service, map[string][]string, bool) { - services, err := transaction.GetAccountServices(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get services for affected group resolution: %v", err) - return nil, nil, false - } - if len(services) == 0 { - return nil, nil, false - } - - proxyByCluster, err := transaction.GetEmbeddedProxyPeerIDsByCluster(ctx, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get embedded proxy peers for affected group resolution: %v", err) - return nil, nil, false - } - if len(proxyByCluster) == 0 { - return nil, nil, false - } - - return services, proxyByCluster, true -} - -func expandChangedPeersWithGroups(ctx context.Context, transaction store.Store, accountID string, changedGroupSet, changedPeerSet map[string]struct{}) map[string]struct{} { - if len(changedGroupSet) == 0 { - return changedPeerSet - } - - ids, err := transaction.GetPeerIDsByGroups(ctx, accountID, setToSlice(changedGroupSet)) - if err != nil { - log.WithContext(ctx).Errorf("failed to expand changed groups to peers for service resolution: %v", err) - return changedPeerSet - } - if len(ids) == 0 { - return changedPeerSet - } - - merged := make(map[string]struct{}, len(changedPeerSet)+len(ids)) - for id := range changedPeerSet { - merged[id] = struct{}{} - } - for _, id := range ids { - merged[id] = struct{}{} - } - return merged -} - -func serviceMatchesChangedPeers(svc *rpservice.Service, proxyPeers []string, changedPeers map[string]struct{}) bool { - for _, pid := range proxyPeers { - if _, ok := changedPeers[pid]; ok { - return true - } - } - for _, target := range svc.Targets { - if target.TargetType != rpservice.TargetTypePeer || target.TargetId == "" { - continue - } - if _, ok := changedPeers[target.TargetId]; ok { - return true - } - } - return false -} - -func addServicePeersToSet(svc *rpservice.Service, proxyPeers []string, peerSet map[string]struct{}) { - for _, pid := range proxyPeers { - peerSet[pid] = struct{}{} - } - for _, target := range svc.Targets { - if target.TargetType == rpservice.TargetTypePeer && target.TargetId != "" { - peerSet[target.TargetId] = struct{}{} - } - } -} - -// collectPolicyRouterBridge folds in the routing peers that serve any network -// resource targeted by the given policies. A policy granting a source access to -// a resource makes the resource's routing peers serve that source at network-map -// compute time, so those routing peers must be refreshed alongside the rule's -// literal groups and peers. The routing peer is reachable only through the -// resource's network, never through the policy's groups, so it is collected here. -// -// Enabled is intentionally not consulted on resources or routers: toggling it is -// itself a change the routing peer must observe. -func collectPolicyRouterBridge(ctx context.Context, transaction store.Store, accountID string, groupSet, peerSet map[string]struct{}, policies ...*types.Policy) { - resourceIDs := policyDestinationResourceIDs(ctx, transaction, accountID, policies...) - if len(resourceIDs) == 0 { - return - } - - networkIDs := resourceNetworkIDs(ctx, transaction, accountID, resourceIDs) - if len(networkIDs) == 0 { - return - } - - routers, err := transaction.GetNetworkRoutersByAccountID(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get network routers for policy router bridge: %v", err) - return - } - - for _, router := range routers { - if _, ok := networkIDs[router.NetworkID]; !ok { - continue - } - addAllToSet(groupSet, router.PeerGroups) - if router.Peer != "" { - peerSet[router.Peer] = struct{}{} - } - } -} - -// policyDestinationResourceIDs collects the network resource IDs targeted by the -// policies' destinations, both via destination groups and via DestinationResource. -func policyDestinationResourceIDs(ctx context.Context, transaction store.Store, accountID string, policies ...*types.Policy) map[string]struct{} { - destGroupSet := make(map[string]struct{}) - resourceIDs := make(map[string]struct{}) - - for _, policy := range policies { - if policy == nil { - continue - } - for _, rule := range policy.Rules { - for _, gID := range rule.Destinations { - destGroupSet[gID] = struct{}{} - } - if rule.DestinationResource.Type != types.ResourceTypePeer && rule.DestinationResource.ID != "" { - resourceIDs[rule.DestinationResource.ID] = struct{}{} - } - } - } - - if len(destGroupSet) > 0 { - groups, err := transaction.GetGroupsByIDs(ctx, store.LockingStrengthNone, accountID, setToSlice(destGroupSet)) - if err != nil { - log.WithContext(ctx).Errorf("failed to get destination groups for policy router bridge: %v", err) - } else { - for _, group := range groups { - for _, res := range group.Resources { - if res.ID != "" { - resourceIDs[res.ID] = struct{}{} - } - } - } - } - } - - return resourceIDs -} - -// resourceNetworkIDs maps the given resource IDs to the set of network IDs that own them. -func resourceNetworkIDs(ctx context.Context, transaction store.Store, accountID string, resourceIDs map[string]struct{}) map[string]struct{} { - resources, err := transaction.GetNetworkResourcesByAccountID(ctx, store.LockingStrengthNone, accountID) - if err != nil { - log.WithContext(ctx).Errorf("failed to get network resources for policy router bridge: %v", err) - return nil - } - - networkIDs := make(map[string]struct{}) - for _, resource := range resources { - if _, ok := resourceIDs[resource.ID]; ok { - networkIDs[resource.NetworkID] = struct{}{} - } - } - return networkIDs -} - -func collectPolicyDirectPeers(policy *types.Policy, peerSet map[string]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{}{} - } - } -} - -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 policyReferencesDirectPeers(policy *types.Policy, changedSet map[string]struct{}) bool { - for _, rule := range policy.Rules { - 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 routeReferencesGroups(r *route.Route, groupSet map[string]struct{}) bool { - return anyInSet(r.Groups, groupSet) || anyInSet(r.PeerGroups, groupSet) || anyInSet(r.AccessControlGroups, groupSet) -} - -func routerReferencesGroups(router *routerTypes.NetworkRouter, groupSet map[string]struct{}) bool { - 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 { - set[id] = struct{}{} - } - } -} - -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/affected_peers_coverage_test.go b/management/server/affected_peers_coverage_test.go new file mode 100644 index 000000000..e62aba90e --- /dev/null +++ b/management/server/affected_peers_coverage_test.go @@ -0,0 +1,119 @@ +package server + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/server/affectedpeers" + "github.com/netbirdio/netbird/management/server/posture" + "github.com/netbirdio/netbird/management/server/types" +) + +// TestAffectedPeers_DependencyCoverageMatrix enumerates each network-map +// dependency crossed with the change-type that can alter it, asserting the +// resolver folds in exactly the peers whose map changes. A new dependency that +// the resolver fails to walk should fail one of these rows; a new change-type +// without a row is a coverage gap to add here. +func TestAffectedPeers_DependencyCoverageMatrix(t *testing.T) { + type row struct { + name string + build func(t *testing.T, s *routerScenario, ctx context.Context) (affectedpeers.Change, []string, []string) + } + + rows := []row{ + { + name: "policy-groups/source-group-change refreshes source+routing, excludes unrelated", + build: func(t *testing.T, s *routerScenario, ctx context.Context) (affectedpeers.Change, []string, []string) { + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + return affectedpeers.Change{ChangedGroupIDs: []string{s.sourceGroupID}}, + []string{s.sourcePeerID, s.routerPeerID}, []string{s.unrelatedPeerID} + }, + }, + { + name: "resource-routing-bridge/router-peer-change refreshes policy sources", + build: func(t *testing.T, s *routerScenario, ctx context.Context) (affectedpeers.Change, []string, []string) { + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + return affectedpeers.Change{ChangedPeerIDs: []string{s.routerPeerID}}, + []string{s.sourcePeerID}, []string{s.unrelatedPeerID} + }, + }, + { + name: "policy-change/explicit-policy refreshes source+routing", + build: func(t *testing.T, s *routerScenario, ctx context.Context) (affectedpeers.Change, []string, []string) { + policy := peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID) + return affectedpeers.Change{Policies: []*types.Policy{policy}}, + []string{s.sourcePeerID, s.routerPeerID}, []string{s.unrelatedPeerID} + }, + }, + { + name: "policy-destinationresource/explicit-policy bridges to routing peer", + build: func(t *testing.T, s *routerScenario, ctx context.Context) (affectedpeers.Change, []string, []string) { + policy := peerToResourcePolicyByResource(s.sourceGroupID, s.resourceID) + return affectedpeers.Change{Policies: []*types.Policy{policy}}, + []string{s.sourcePeerID, s.routerPeerID}, []string{s.unrelatedPeerID} + }, + }, + { + name: "resource-change refreshes source+routing on its network", + build: func(t *testing.T, s *routerScenario, ctx context.Context) (affectedpeers.Change, []string, []string) { + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + return affectedpeers.Change{ResourceIDs: []string{s.resourceID}}, + []string{s.sourcePeerID, s.routerPeerID}, []string{s.unrelatedPeerID} + }, + }, + { + name: "network-change refreshes source+routing on that network", + build: func(t *testing.T, s *routerScenario, ctx context.Context) (affectedpeers.Change, []string, []string) { + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + return affectedpeers.Change{NetworkIDs: []string{s.networkID}}, + []string{s.sourcePeerID, s.routerPeerID}, []string{s.unrelatedPeerID} + }, + }, + { + name: "posture-check-change refreshes source+routing of gated policy", + build: func(t *testing.T, s *routerScenario, ctx context.Context) (affectedpeers.Change, []string, []string) { + check, err := s.manager.SavePostureChecks(ctx, s.accountID, userID, &posture.Checks{ + Name: "cov-min-version", + Checks: posture.ChecksDefinition{NBVersionCheck: &posture.NBVersionCheck{MinVersion: "0.30.0"}}, + }, true) + require.NoError(t, err) + policy := peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID) + policy.SourcePostureChecks = []string{check.ID} + _, err = s.manager.SavePolicy(ctx, s.accountID, userID, policy, true) + require.NoError(t, err) + return affectedpeers.Change{PostureCheckIDs: []string{check.ID}}, + []string{s.sourcePeerID, s.routerPeerID}, []string{s.unrelatedPeerID} + }, + }, + { + name: "empty-change yields nothing", + build: func(t *testing.T, s *routerScenario, ctx context.Context) (affectedpeers.Change, []string, []string) { + return affectedpeers.Change{}, nil, []string{s.sourcePeerID, s.routerPeerID, s.unrelatedPeerID} + }, + }, + } + + for _, r := range rows { + t.Run(r.name, func(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + change, mustContain, mustExclude := r.build(t, s, ctx) + affected := s.manager.ResolveAffectedPeers(ctx, s.manager.Store, s.accountID, change) + + for _, id := range mustContain { + assert.Contains(t, affected, id, "expected peer to be affected") + } + for _, id := range mustExclude { + assert.NotContains(t, affected, id, "peer must not be affected") + } + }) + } +} diff --git a/management/server/affected_peers_oldstate_test.go b/management/server/affected_peers_oldstate_test.go new file mode 100644 index 000000000..bcb78a660 --- /dev/null +++ b/management/server/affected_peers_oldstate_test.go @@ -0,0 +1,143 @@ +package server + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/require" + + resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" + routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" +) + +// An update spans an old and a new state. The affected set must be the UNION of +// peers reachable before and after the change; resolving only against the final +// state drops peers that were reachable but no longer are. These tests pin the +// two paths where the old state is reachable only by the changed object's +// previous references: detaching a resource group, and re-pointing a router peer. + +// TestAffectedPeers_E2E_UpdateResource_DetachGroup_RefreshesOldGroupSources: +// a resource is reachable by a source group via two destination resource groups; +// detaching one of them must still refresh that group's policy source peers, even +// though the post-update resource no longer maps to it. +func TestAffectedPeers_E2E_UpdateResource_DetachGroup_RefreshesOldGroupSources(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + // A second resource group + a second source group/peer that reaches the + // resource only through that second group. + const detachGroupID = "rs-detach-grp" + require.NoError(t, s.manager.CreateGroup(ctx, s.accountID, userID, &types.Group{ID: detachGroupID, Name: "rs-detach"})) + + const secondSourceGroupID = "rs-source-grp-2" + setupKey, err := s.manager.CreateSetupKey(ctx, s.accountID, "rs-detach-key", types.SetupKeyReusable, time.Hour, nil, 999, userID, false, false) + require.NoError(t, err) + secondSourcePeer := addPeerToAccount(t, s.manager, s.accountID, setupKey.Key) + require.NoError(t, s.manager.CreateGroup(ctx, s.accountID, userID, &types.Group{ + ID: secondSourceGroupID, Name: "rs-source-2", Peers: []string{secondSourcePeer.ID}, + })) + + resourcesManager, _, _ := s.managers() + + // Attach the resource to the detach group as well: now in [resourceGroup, detachGroup]. + _, err = resourcesManager.UpdateResource(ctx, userID, &resourceTypes.NetworkResource{ + ID: s.resourceID, + AccountID: s.accountID, + NetworkID: s.networkID, + Name: "rs-resource-host", + Address: "10.20.30.0/24", + GroupIDs: []string{s.resourceGroupID, detachGroupID}, + Enabled: true, + }) + require.NoError(t, err) + + // Policy granting the second source group access via the detach group. + _, err = s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(secondSourceGroupID, detachGroupID), true) + require.NoError(t, err) + + secondSrcCh := s.updateManager.CreateChannel(ctx, secondSourcePeer.ID) + t.Cleanup(func() { s.updateManager.CloseChannel(ctx, secondSourcePeer.ID) }) + settleAffectedUpdates(secondSrcCh) + + done := make(chan struct{}) + go func() { + // Detaching the resource from detachGroup removes the second source's + // access; that source peer must be refreshed even though the post-update + // resource no longer maps to detachGroup. + peerShouldReceiveUpdate(t, secondSrcCh) + close(done) + }() + + _, err = resourcesManager.UpdateResource(ctx, userID, &resourceTypes.NetworkResource{ + ID: s.resourceID, + AccountID: s.accountID, + NetworkID: s.networkID, + Name: "rs-resource-host", + Address: "10.20.30.0/24", + GroupIDs: []string{s.resourceGroupID}, // detached detachGroup + Enabled: true, + }) + require.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout: detaching a resource group did not refresh the old group's policy source peer") + } +} + +// TestAffectedPeers_E2E_UpdateRouter_RepointPeer_RefreshesOldRoutingPeer: +// changing router.Peer within the same network must still refresh the OLD routing +// peer, which loses its routing role. +func TestAffectedPeers_E2E_UpdateRouter_RepointPeer_RefreshesOldRoutingPeer(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + + _, routersManager, _ := s.managers() + + routers, err := s.manager.Store.GetNetworkRoutersByNetID(ctx, store.LockingStrengthNone, s.accountID, s.networkID) + require.NoError(t, err) + require.Len(t, routers, 1) + router := routers[0] + oldRoutingPeer := router.Peer + require.NotEmpty(t, oldRoutingPeer) + + // A new peer to become the routing peer in place of the old one. + setupKey, err := s.manager.CreateSetupKey(ctx, s.accountID, "rs-newrouter-key", types.SetupKeyReusable, time.Hour, nil, 999, userID, false, false) + require.NoError(t, err) + newRoutingPeer := addPeerToAccount(t, s.manager, s.accountID, setupKey.Key) + + oldCh := s.updateManager.CreateChannel(ctx, oldRoutingPeer) + t.Cleanup(func() { s.updateManager.CloseChannel(ctx, oldRoutingPeer) }) + settleAffectedUpdates(oldCh) + + done := make(chan struct{}) + go func() { + // The old routing peer stops serving the resource and must be refreshed. + peerShouldReceiveUpdate(t, oldCh) + close(done) + }() + + _, err = routersManager.UpdateRouter(ctx, userID, &routerTypes.NetworkRouter{ + ID: router.ID, + NetworkID: s.networkID, + AccountID: s.accountID, + Peer: newRoutingPeer.ID, // repoint within the same network + Masquerade: true, + Metric: 9999, + Enabled: true, + }) + require.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout: re-pointing the router peer did not refresh the old routing peer") + } +} diff --git a/management/server/affected_peers_property_test.go b/management/server/affected_peers_property_test.go new file mode 100644 index 000000000..82dd91cf4 --- /dev/null +++ b/management/server/affected_peers_property_test.go @@ -0,0 +1,251 @@ +package server + +import ( + "context" + "encoding/json" + "fmt" + "math/rand" + "sort" + "testing" + + "github.com/stretchr/testify/require" + "golang.org/x/exp/maps" + + nbdns "github.com/netbirdio/netbird/dns" + "github.com/netbirdio/netbird/management/server/affectedpeers" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" +) + +// allPeerMaps computes the serialized per-peer network map for every peer in the +// account, mirroring the controller's compute path so the property test compares +// against real output. +func allPeerMaps(t *testing.T, manager *DefaultAccountManager, accountID string) map[string]string { + t.Helper() + ctx := context.Background() + + account, err := manager.Store.GetAccount(ctx, accountID) + require.NoError(t, err) + + account.InjectProxyPolicies(ctx) + + validated := make(map[string]struct{}, len(account.Peers)) + for id := range account.Peers { + validated[id] = struct{}{} + } + resourcePolicies := account.GetResourcePoliciesMap() + routers := account.GetResourceRoutersMap() + groupIDToUserIDs := account.GetActiveGroupUsers() + + out := make(map[string]string, len(account.Peers)) + for peerID := range account.Peers { + nm := account.GetPeerNetworkMapFromComponents(ctx, peerID, nbdns.CustomZone{}, nil, validated, resourcePolicies, routers, nil, groupIDToUserIDs) + // Network.Serial is an account-global counter bumped on every change; it + // is not a per-peer dependency, so normalize it out of the comparison. + if nm.Network != nil { + nm.Network.Serial = 0 + } + out[peerID] = canonicalJSON(t, nm) + } + return out +} + +// canonicalJSON marshals v and returns an order-insensitive string form: every +// JSON array is sorted by the canonical form of its elements. The network map's +// Peers/Routes/FirewallRules/SourceRanges slices have nondeterministic order, so +// a raw JSON compare would report spurious changes. +func canonicalJSON(t *testing.T, v interface{}) string { + t.Helper() + b, err := json.Marshal(v) + require.NoError(t, err) + var parsed interface{} + require.NoError(t, json.Unmarshal(b, &parsed)) + canonicalized, err := json.Marshal(sortAny(parsed)) + require.NoError(t, err) + return string(canonicalized) +} + +func sortAny(v interface{}) interface{} { + switch val := v.(type) { + case []interface{}: + for i := range val { + val[i] = sortAny(val[i]) + } + sort.Slice(val, func(i, j int) bool { + bi, _ := json.Marshal(val[i]) + bj, _ := json.Marshal(val[j]) + return string(bi) < string(bj) + }) + return val + case map[string]interface{}: + for k := range val { + val[k] = sortAny(val[k]) + } + return val + default: + return v + } +} + +// changedPeers returns the peer IDs whose serialized map differs between before +// and after. +func changedPeers(before, after map[string]string) []string { + var changed []string + for id, b := range before { + a, ok := after[id] + if !ok || a != b { + changed = append(changed, id) + } + } + for id := range after { + if _, ok := before[id]; !ok { + changed = append(changed, id) + } + } + return changed +} + +// TestAffectedPeers_Property_ResolverSupersetsRealChanges builds a topology, +// applies random changes, and asserts that the resolver's affected set is a +// superset of the peers whose real network map actually changed. If the resolver +// ever misses a dependency, a change will alter a peer's map without that peer +// appearing in the affected set, failing here. +func TestAffectedPeers_Property_ResolverSupersetsRealChanges(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + // A pre-existing peer->resource policy so the resource/router bridge is live. + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + + // Extra peers and groups to give mutations room to move membership around. + setupKey, err := s.manager.CreateSetupKey(ctx, s.accountID, "prop-key", types.SetupKeyReusable, 0, nil, 999, userID, false, false) + require.NoError(t, err) + extraPeers := make([]string, 0, 4) + for i := 0; i < 4; i++ { + p := addPeerToAccount(t, s.manager, s.accountID, setupKey.Key) + extraPeers = append(extraPeers, p.ID) + } + extraGroups := []string{"prop-grp-0", "prop-grp-1"} + for _, g := range extraGroups { + require.NoError(t, s.manager.CreateGroup(ctx, s.accountID, userID, &types.Group{ID: g, Name: g})) + } + + rng := rand.New(rand.NewSource(1)) + allGroups := append([]string{s.sourceGroupID, s.resourceGroupID, s.routerPeerGroupID}, extraGroups...) + allPeers := append([]string{s.sourcePeerID, s.routerPeerID, s.routerGroupPeerID, s.unrelatedPeerID}, extraPeers...) + + for iter := 0; iter < 60; iter++ { + change, apply := s.randomMutation(t, rng, allGroups, allPeers) + if apply == nil { + continue + } + + before := allPeerMaps(t, s.manager, s.accountID) + + resolvedSet := make(map[string]struct{}) + resolve := func() { + require.NoError(t, s.manager.Store.ExecuteInTransaction(ctx, func(tx store.Store) error { + for _, id := range s.manager.ResolveAffectedPeers(ctx, tx, s.accountID, change) { + resolvedSet[id] = struct{}{} + } + return nil + })) + } + + // Resolve on both sides of the mutation and union: removals are visible + // only pre-apply (the leaving peer is still a member), additions only + // post-apply (the joining peer is now a member). Production captures both + // via per-path handling (e.g. UpdateGroup passes peersToRemove); the union + // models that without coupling the test to each path's ordering. + resolve() + changedIDs := change.ChangedPeerIDs + apply() + resolve() + + after := allPeerMaps(t, s.manager, s.accountID) + + // The explicitly-changed peer's own map refresh is the caller's + // responsibility (the resolver returns the peers to propagate to), so it + // is allowed to be absent from the resolved set. + changedExplicitly := make(map[string]struct{}, len(changedIDs)) + for _, id := range changedIDs { + changedExplicitly[id] = struct{}{} + } + + for _, id := range changedPeers(before, after) { + if _, stillExists := after[id]; !stillExists { + continue + } + if _, isExplicit := changedExplicitly[id]; isExplicit { + continue + } + _, ok := resolvedSet[id] + require.Truef(t, ok, + "iter %d: peer %s network map changed but was not in the resolver's affected set %v (change=%+v)", + iter, id, maps.Keys(resolvedSet), change) + } + } +} + +// randomMutation picks a random change, returns the Change to resolve and a +// function that applies the underlying store mutation. apply is nil when the +// drawn mutation is a no-op for the current state. +func (s *routerScenario) randomMutation(t *testing.T, rng *rand.Rand, allGroups, allPeers []string) (affectedpeers.Change, func()) { + t.Helper() + ctx := context.Background() + + switch rng.Intn(3) { + case 0: + groupID := allGroups[rng.Intn(len(allGroups))] + peerID := allPeers[rng.Intn(len(allPeers))] + grp, err := s.manager.Store.GetGroupByID(ctx, store.LockingStrengthNone, s.accountID, groupID) + require.NoError(t, err) + if slicesContains(grp.Peers, peerID) { + return affectedpeers.Change{}, nil + } + return affectedpeers.Change{ChangedGroupIDs: []string{groupID}, ChangedPeerIDs: []string{peerID}}, + func() { + require.NoError(t, s.manager.GroupAddPeer(ctx, s.accountID, groupID, peerID)) + } + case 1: + groupID := allGroups[rng.Intn(len(allGroups))] + grp, err := s.manager.Store.GetGroupByID(ctx, store.LockingStrengthNone, s.accountID, groupID) + require.NoError(t, err) + if len(grp.Peers) == 0 { + return affectedpeers.Change{}, nil + } + peerID := grp.Peers[rng.Intn(len(grp.Peers))] + return affectedpeers.Change{ChangedGroupIDs: []string{groupID}, ChangedPeerIDs: []string{peerID}}, + func() { + require.NoError(t, s.manager.GroupDeletePeer(ctx, s.accountID, groupID, peerID)) + } + default: + src := allGroups[rng.Intn(len(allGroups))] + dst := allGroups[rng.Intn(len(allGroups))] + policy := &types.Policy{ + Enabled: true, + Name: fmt.Sprintf("prop-policy-%d", rng.Int()), + Rules: []*types.PolicyRule{{ + Enabled: true, + Sources: []string{src}, + Destinations: []string{dst}, + Action: types.PolicyTrafficActionAccept, + }}, + } + return affectedpeers.Change{Policies: []*types.Policy{policy}}, + func() { + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, policy, true) + require.NoError(t, err) + } + } +} + +func slicesContains(s []string, v string) bool { + for _, x := range s { + if x == v { + return true + } + } + return false +} diff --git a/management/server/affected_peers_querycount_test.go b/management/server/affected_peers_querycount_test.go new file mode 100644 index 000000000..f37be1243 --- /dev/null +++ b/management/server/affected_peers_querycount_test.go @@ -0,0 +1,125 @@ +package server + +import ( + "context" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + nbdns "github.com/netbirdio/netbird/dns" + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/management/server/affectedpeers" + resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" + 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" +) + +// countingStore wraps a real store and counts the per-account collection loads +// the resolver performs, so a test can assert each is read at most once and that +// irrelevant collections are skipped entirely. +type countingStore struct { + store.Store + mu sync.Mutex + counts map[string]int +} + +func newCountingStore(s store.Store) *countingStore { + return &countingStore{Store: s, counts: map[string]int{}} +} + +func (c *countingStore) bump(name string) { + c.mu.Lock() + c.counts[name]++ + c.mu.Unlock() +} + +func (c *countingStore) count(name string) int { + c.mu.Lock() + defer c.mu.Unlock() + return c.counts[name] +} + +func (c *countingStore) GetAccountPolicies(ctx context.Context, ls store.LockingStrength, accountID string) ([]*types.Policy, error) { + c.bump("policies") + return c.Store.GetAccountPolicies(ctx, ls, accountID) +} + +func (c *countingStore) GetAccountRoutes(ctx context.Context, ls store.LockingStrength, accountID string) ([]*route.Route, error) { + c.bump("routes") + return c.Store.GetAccountRoutes(ctx, ls, accountID) +} + +func (c *countingStore) GetAccountNameServerGroups(ctx context.Context, ls store.LockingStrength, accountID string) ([]*nbdns.NameServerGroup, error) { + c.bump("nameservers") + return c.Store.GetAccountNameServerGroups(ctx, ls, accountID) +} + +func (c *countingStore) GetAccountDNSSettings(ctx context.Context, ls store.LockingStrength, accountID string) (*types.DNSSettings, error) { + c.bump("dnssettings") + return c.Store.GetAccountDNSSettings(ctx, ls, accountID) +} + +func (c *countingStore) GetNetworkRoutersByAccountID(ctx context.Context, ls store.LockingStrength, accountID string) ([]*routerTypes.NetworkRouter, error) { + c.bump("routers") + return c.Store.GetNetworkRoutersByAccountID(ctx, ls, accountID) +} + +func (c *countingStore) GetNetworkResourcesByAccountID(ctx context.Context, ls store.LockingStrength, accountID string) ([]*resourceTypes.NetworkResource, error) { + c.bump("resources") + return c.Store.GetNetworkResourcesByAccountID(ctx, ls, accountID) +} + +func (c *countingStore) GetAccountServices(ctx context.Context, ls store.LockingStrength, accountID string) ([]*rpservice.Service, error) { + c.bump("services") + return c.Store.GetAccountServices(ctx, ls, accountID) +} + +// TestAffectedPeers_QueryCount_NoRedundantFullTableLoads asserts the resolver +// loads each per-account collection at most once per Resolve (memoization) even +// on a change that drives every bridge, and skips the services table when the +// account has no embedded proxy peers. +func TestAffectedPeers_QueryCount_NoRedundantFullTableLoads(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + + cs := newCountingStore(s.manager.Store) + + // A group change that exercises policies, routers, resources and the bridge. + affected, err := affectedpeers.Resolve(ctx, cs, s.accountID, affectedpeers.Change{ChangedGroupIDs: []string{s.sourceGroupID}}) + require.NoError(t, err) + assert.Contains(t, affected, s.routerPeerID, "bridge must still resolve the routing peer") + + for _, name := range []string{"policies", "routes", "nameservers", "dnssettings", "routers", "resources"} { + assert.LessOrEqualf(t, cs.count(name), 1, + "%s must be loaded at most once per Resolve, got %d", name, cs.count(name)) + } + assert.Equal(t, 0, cs.count("services"), + "services must not be loaded when the account has no embedded proxy peers") +} + +// TestAffectedPeers_QueryCount_NarrowChangeSkipsLoads asserts that a change with +// no group/peer signal touches no per-account collections beyond what its inputs +// require. +func TestAffectedPeers_QueryCount_NarrowChangeSkipsLoads(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + cs := newCountingStore(s.manager.Store) + + // A bare network-id change drives only the router->source bridge: routers and + // resources are needed, but routes/nameservers/dnssettings/services are not. + _, err := affectedpeers.Resolve(ctx, cs, s.accountID, affectedpeers.Change{NetworkIDs: []string{s.networkID}}) + require.NoError(t, err) + + assert.Equal(t, 0, cs.count("routes"), "routes must not be loaded for a network-only change") + assert.Equal(t, 0, cs.count("nameservers"), "nameservers must not be loaded for a network-only change") + assert.Equal(t, 0, cs.count("dnssettings"), "dnssettings must not be loaded for a network-only change") + assert.Equal(t, 0, cs.count("services"), "services must not be loaded for a network-only change") +} diff --git a/management/server/affected_peers_router_paths_test.go b/management/server/affected_peers_router_paths_test.go new file mode 100644 index 000000000..bdd7a6cc4 --- /dev/null +++ b/management/server/affected_peers_router_paths_test.go @@ -0,0 +1,369 @@ +package server + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/server/affectedpeers" + resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" + routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" + "github.com/netbirdio/netbird/management/server/posture" + "github.com/netbirdio/netbird/management/server/types" +) + +func (s *routerScenario) resolveGroupChangeAffected(ctx context.Context, changedGroupIDs []string) []string { + return s.manager.ResolveAffectedPeers(ctx, s.manager.Store, s.accountID, affectedpeers.Change{ChangedGroupIDs: changedGroupIDs}) +} + +func (s *routerScenario) resolvePeerChangeAffected(ctx context.Context, changedPeerIDs []string) []string { + return s.manager.resolveAffectedPeersForPeerChanges(ctx, s.manager.Store, s.accountID, changedPeerIDs) +} + +func TestAffectedPeers_GroupChange_SourceGroupMembership_RefreshesRoutingPeer_DirectRouter(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + + affected := s.resolveGroupChangeAffected(ctx, []string{s.sourceGroupID}) + + assert.Contains(t, affected, s.sourcePeerID, "source group member must be affected") + assert.Contains(t, affected, s.routerPeerID, + "changing the source group of a peer->resource policy must refresh the resource's routing peer") + assert.NotContains(t, affected, s.unrelatedPeerID, "unrelated peer must not be affected") +} + +func TestAffectedPeers_GroupChange_SourceGroupMembership_RefreshesRoutingPeer_RouterPeerGroups(t *testing.T) { + s := setupRouterScenario(t, false) + ctx := context.Background() + + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + + affected := s.resolveGroupChangeAffected(ctx, []string{s.sourceGroupID}) + + assert.Contains(t, affected, s.routerGroupPeerID, + "changing the source group must refresh the routing peer defined via router.PeerGroups") + assert.NotContains(t, affected, s.unrelatedPeerID, "unrelated peer must not be affected") +} + +func TestAffectedPeers_GroupChange_RouterPeerGroupMembership_RefreshesPolicySources(t *testing.T) { + s := setupRouterScenario(t, false) + ctx := context.Background() + + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + + affected := s.resolveGroupChangeAffected(ctx, []string{s.routerPeerGroupID}) + + assert.Contains(t, affected, s.routerGroupPeerID, "the routing peer itself must be affected") + assert.Contains(t, affected, s.sourcePeerID, + "changing the router's PeerGroups must refresh the source peers of policies serving the resource") + assert.NotContains(t, affected, s.unrelatedPeerID, "unrelated peer must not be affected") +} + +func TestAffectedPeers_PeerChange_SourcePeer_RefreshesRoutingPeer(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + + affected := s.resolvePeerChangeAffected(ctx, []string{s.sourcePeerID}) + + assert.Contains(t, affected, s.routerPeerID, + "a status change on a source peer must refresh the resource's routing peer that serves it") + assert.NotContains(t, affected, s.unrelatedPeerID, "unrelated peer must not be affected") +} + +func TestAffectedPeers_PeerChange_RoutingPeer_RefreshesPolicySources(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + + affected := s.resolvePeerChangeAffected(ctx, []string{s.routerPeerID}) + + assert.Contains(t, affected, s.sourcePeerID, + "a status change on the routing peer must refresh the source peers that route through it") + assert.NotContains(t, affected, s.unrelatedPeerID, "unrelated peer must not be affected") +} + +func TestAffectedPeers_PeerChange_SourcePeer_ByDestinationResource_RefreshesRoutingPeer(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByResource(s.sourceGroupID, s.resourceID), true) + require.NoError(t, err) + + affected := s.resolvePeerChangeAffected(ctx, []string{s.sourcePeerID}) + + assert.Contains(t, affected, s.routerPeerID, + "DestinationResource-targeted policy must still bridge a source-peer change to the routing peer") + assert.NotContains(t, affected, s.unrelatedPeerID, "unrelated peer must not be affected") +} + +func TestAffectedPeers_E2E_DeleteGroup_ResolvesAffectedPeers(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + const memberOnlyGroupID = "rs-memberonly-grp" + require.NoError(t, s.manager.CreateGroup(ctx, s.accountID, userID, &types.Group{ + ID: memberOnlyGroupID, Name: "rs-memberonly", Peers: []string{s.sourcePeerID}, + })) + + affected := s.resolveGroupChangeAffected(ctx, []string{memberOnlyGroupID}) + assert.Empty(t, affected, "an unlinked group has no network-map impact, so no peer is affected") + + require.NoError(t, s.manager.DeleteGroup(ctx, s.accountID, userID, memberOnlyGroupID)) +} + +func TestAffectedPeers_DeleteGroup_LinkedGroupIsBlocked(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + + err = s.manager.DeleteGroup(ctx, s.accountID, userID, s.sourceGroupID) + require.Error(t, err, "deleting a policy-linked group must be blocked by validateDeleteGroup") + + var linkErr *GroupLinkError + require.ErrorAs(t, err, &linkErr, "expected a GroupLinkError") +} + +func TestAffectedPeers_GroupAddResource_RefreshesRoutingPeer(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + const extraResourceGroupID = "rs-resource-grp-extra" + require.NoError(t, s.manager.CreateGroup(ctx, s.accountID, userID, &types.Group{ + ID: extraResourceGroupID, Name: "rs-resource-extra", + })) + + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, extraResourceGroupID), true) + require.NoError(t, err) + + require.NoError(t, s.manager.GroupAddResource(ctx, s.accountID, extraResourceGroupID, types.Resource{ + ID: s.resourceID, + Type: types.ResourceTypeHost, + })) + + affected := s.resolveGroupChangeAffected(ctx, []string{extraResourceGroupID}) + + assert.Contains(t, affected, s.routerPeerID, + "attaching a resource to a policy destination group must refresh the resource's routing peer") + assert.Contains(t, affected, s.sourcePeerID, "policy source peers must refresh") + assert.NotContains(t, affected, s.unrelatedPeerID, "unrelated peer must not be affected") +} + +func (s *routerScenario) resolvePostureCheckAffected(ctx context.Context, postureCheckID string) []string { + return s.manager.ResolveAffectedPeers(ctx, s.manager.Store, s.accountID, affectedpeers.Change{PostureCheckIDs: []string{postureCheckID}}) +} + +func (s *routerScenario) createPostureCheckGatedPolicy(t *testing.T, ctx context.Context) string { + t.Helper() + + check, err := s.manager.SavePostureChecks(ctx, s.accountID, userID, &posture.Checks{ + Name: "rs-min-version", + Checks: posture.ChecksDefinition{ + NBVersionCheck: &posture.NBVersionCheck{MinVersion: "0.30.0"}, + }, + }, true) + require.NoError(t, err) + + policy := peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID) + policy.SourcePostureChecks = []string{check.ID} + _, err = s.manager.SavePolicy(ctx, s.accountID, userID, policy, true) + require.NoError(t, err) + + return check.ID +} + +func TestAffectedPeers_PostureCheckChange_RefreshesRoutingPeer(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + checkID := s.createPostureCheckGatedPolicy(t, ctx) + + affected := s.resolvePostureCheckAffected(ctx, checkID) + + assert.Contains(t, affected, s.sourcePeerID, "policy source peer must be affected by a posture-check change") + assert.Contains(t, affected, s.routerPeerID, + "a posture check gating a peer->resource policy must refresh the resource's routing peer") + assert.NotContains(t, affected, s.unrelatedPeerID, "unrelated peer must not be affected") +} + +func TestAffectedPeers_E2E_SavePostureCheck_RefreshesRoutingPeer(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + checkID := s.createPostureCheckGatedPolicy(t, ctx) + + srcCh := s.updateManager.CreateChannel(ctx, s.sourcePeerID) + routerCh := s.updateManager.CreateChannel(ctx, s.routerPeerID) + unrelatedCh := s.updateManager.CreateChannel(ctx, s.unrelatedPeerID) + t.Cleanup(func() { + s.updateManager.CloseChannel(ctx, s.sourcePeerID) + s.updateManager.CloseChannel(ctx, s.routerPeerID) + s.updateManager.CloseChannel(ctx, s.unrelatedPeerID) + }) + + settleAffectedUpdates(srcCh, routerCh, unrelatedCh) + + done := make(chan struct{}) + go func() { + peerShouldReceiveUpdate(t, srcCh) + peerShouldReceiveUpdate(t, routerCh) + peerShouldNotReceiveUpdate(t, unrelatedCh) + close(done) + }() + + _, err := s.manager.SavePostureChecks(ctx, s.accountID, userID, &posture.Checks{ + ID: checkID, + Name: "rs-min-version", + Checks: posture.ChecksDefinition{ + NBVersionCheck: &posture.NBVersionCheck{MinVersion: "0.31.0"}, + }, + }, false) + require.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout: editing a posture check did not refresh source + routing peers") + } +} + +func TestAffectedPeers_E2E_UpdateResource_DestinationResourcePolicy_RefreshesSourcePeer(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByResource(s.sourceGroupID, s.resourceID), true) + require.NoError(t, err) + + resourcesManager, _, _ := s.managers() + + srcCh := s.updateManager.CreateChannel(ctx, s.sourcePeerID) + routerCh := s.updateManager.CreateChannel(ctx, s.routerPeerID) + unrelatedCh := s.updateManager.CreateChannel(ctx, s.unrelatedPeerID) + t.Cleanup(func() { + s.updateManager.CloseChannel(ctx, s.sourcePeerID) + s.updateManager.CloseChannel(ctx, s.routerPeerID) + s.updateManager.CloseChannel(ctx, s.unrelatedPeerID) + }) + + settleAffectedUpdates(srcCh, routerCh, unrelatedCh) + + done := make(chan struct{}) + go func() { + peerShouldReceiveUpdate(t, srcCh) + peerShouldReceiveUpdate(t, routerCh) + peerShouldNotReceiveUpdate(t, unrelatedCh) + close(done) + }() + + _, err = resourcesManager.UpdateResource(ctx, userID, &resourceTypes.NetworkResource{ + ID: s.resourceID, + AccountID: s.accountID, + NetworkID: s.networkID, + Name: "rs-resource-host", + Address: "10.20.30.0/25", + GroupIDs: []string{s.resourceGroupID}, + Enabled: true, + }) + require.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout: updating a DestinationResource-targeted resource did not refresh its policy source peer") + } +} + +func TestAffectedPeers_E2E_UpdateResource_DisabledSiblingRouter_StillBridged(t *testing.T) { + s := setupRouterScenario(t, true) + ctx := context.Background() + + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + + resourcesManager, routersManager, _ := s.managers() + + setupKey, err := s.manager.CreateSetupKey(ctx, s.accountID, "rs-key-disabled", types.SetupKeyReusable, time.Hour, nil, 999, userID, false, false) + require.NoError(t, err) + disabledRouterPeer := addPeerToAccount(t, s.manager, s.accountID, setupKey.Key) + _, err = routersManager.CreateRouter(ctx, userID, &routerTypes.NetworkRouter{ + NetworkID: s.networkID, + AccountID: s.accountID, + Peer: disabledRouterPeer.ID, + Masquerade: true, + Metric: 9000, + Enabled: false, + }) + require.NoError(t, err) + + disabledCh := s.updateManager.CreateChannel(ctx, disabledRouterPeer.ID) + t.Cleanup(func() { s.updateManager.CloseChannel(ctx, disabledRouterPeer.ID) }) + + settleAffectedUpdates(disabledCh) + + done := make(chan struct{}) + go func() { + peerShouldReceiveUpdate(t, disabledCh) + close(done) + }() + + _, err = resourcesManager.UpdateResource(ctx, userID, &resourceTypes.NetworkResource{ + ID: s.resourceID, + AccountID: s.accountID, + NetworkID: s.networkID, + Name: "rs-resource-host", + Address: "10.20.30.0/25", + GroupIDs: []string{s.resourceGroupID}, + Enabled: true, + }) + require.NoError(t, err) + + select { + case <-done: + case <-time.After(peerUpdateTimeout): + t.Error("timeout: resource update did not refresh the disabled sibling router's peer") + } +} + +func TestAffectedPeers_GroupChange_RouterInOtherNetworkNotAffected(t *testing.T) { + s := setupRouterScenario(t, true) + second := s.addSecondTopology(t, "groupiso") + ctx := context.Background() + + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + + affected := s.resolveGroupChangeAffected(ctx, []string{s.sourceGroupID}) + + assert.Contains(t, affected, s.routerPeerID, "network A's routing peer must be affected") + assert.NotContains(t, affected, second.routerPeerID, + "a router in an unrelated network must not be affected by a source-group change for another resource") +} + +func TestAffectedPeers_PeerChange_RouterInOtherNetworkNotAffected(t *testing.T) { + s := setupRouterScenario(t, true) + second := s.addSecondTopology(t, "peeriso") + ctx := context.Background() + + _, err := s.manager.SavePolicy(ctx, s.accountID, userID, peerToResourcePolicyByGroup(s.sourceGroupID, s.resourceGroupID), true) + require.NoError(t, err) + + affected := s.resolvePeerChangeAffected(ctx, []string{s.sourcePeerID}) + + assert.Contains(t, affected, s.routerPeerID, "network A's routing peer must be affected") + assert.NotContains(t, affected, second.routerPeerID, + "a router in an unrelated network must not be affected by a source-peer change for another resource") +} diff --git a/management/server/affected_peers_router_test.go b/management/server/affected_peers_router_test.go index 77e31104a..099f32533 100644 --- a/management/server/affected_peers_router_test.go +++ b/management/server/affected_peers_router_test.go @@ -10,6 +10,7 @@ import ( "github.com/netbirdio/netbird/management/internals/controllers/network_map" "github.com/netbirdio/netbird/management/internals/controllers/network_map/update_channel" + "github.com/netbirdio/netbird/management/server/affectedpeers" "github.com/netbirdio/netbird/management/server/groups" "github.com/netbirdio/netbird/management/server/networks" "github.com/netbirdio/netbird/management/server/networks/resources" @@ -199,10 +200,10 @@ func peerToResourcePolicyByResource(sourceGroupID, resourceID string) *types.Pol // routing peer is expected to be in the affected set but is not. // --------------------------------------------------------------------------- -// resolvePolicyAffected mirrors SavePolicy's resolution: collect groups/peers -// from the policy, then expand to concrete peer IDs. +// resolvePolicyAffected mirrors SavePolicy's resolution: resolve the affected +// peers for the given policy. func (s *routerScenario) resolvePolicyAffected(ctx context.Context, policy *types.Policy) []string { - return s.manager.resolvePolicyAffectedPeers(ctx, s.manager.Store, s.accountID, policy) + return s.manager.ResolveAffectedPeers(ctx, s.manager.Store, s.accountID, affectedpeers.Change{Policies: []*types.Policy{policy}}) } func TestAffectedPeers_PolicyToResourceByGroup_IncludesSourcePeer_DirectRouter(t *testing.T) { diff --git a/management/server/affected_peers_test.go b/management/server/affected_peers_test.go index 1de9add53..ec7f84b12 100644 --- a/management/server/affected_peers_test.go +++ b/management/server/affected_peers_test.go @@ -13,6 +13,7 @@ import ( nbdns "github.com/netbirdio/netbird/dns" rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/management/server/affectedpeers" 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" @@ -22,6 +23,20 @@ import ( "github.com/netbirdio/netbird/route" ) +// Thin test adapters over affectedpeers.Collect, preserving the (groups, peers) +// shape these tests assert on after the resolver was unified. +func collectGroupChangeAffectedGroups(ctx context.Context, s store.Store, accountID string, changedGroupIDs []string) ([]string, []string) { + return affectedpeers.Collect(ctx, s, accountID, affectedpeers.Change{ChangedGroupIDs: changedGroupIDs}) +} + +func collectPeerChangeAffectedGroups(ctx context.Context, s store.Store, accountID string, changedGroupIDs, changedPeerIDs []string) ([]string, []string) { + return affectedpeers.Collect(ctx, s, accountID, affectedpeers.Change{ChangedGroupIDs: changedGroupIDs, ChangedPeerIDs: changedPeerIDs}) +} + +func collectPostureCheckAffectedGroupsAndPeers(ctx context.Context, s store.Store, accountID, postureCheckID string) ([]string, []string) { + return affectedpeers.Collect(ctx, s, accountID, affectedpeers.Change{PostureCheckIDs: []string{postureCheckID}}) +} + // 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) { @@ -425,219 +440,8 @@ func TestCollectGroupChange_MultipleNameServerGroups_OnlyLinkedAffected(t *testi assert.Empty(t, groups) } -// --------------------------------------------------------------------------- -// collectPolicyAffectedGroupsAndPeers unit tests -// --------------------------------------------------------------------------- - -func TestCollectPolicyAffectedGroups_Basic(t *testing.T) { - policy := &types.Policy{ - Rules: []*types.PolicyRule{ - { - Sources: []string{"g1", "g2"}, - Destinations: []string{"g3"}, - }, - }, - } - groups, directPeers := collectPolicyAffectedGroupsAndPeers(context.Background(), 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(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(context.Background(), 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(context.Background(), policy) - 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"}}, - }, - } - updated := &types.Policy{ - Rules: []*types.PolicyRule{ - {Sources: []string{"g3"}, Destinations: []string{"g4"}}, - }, - } - groups, _ := collectPolicyAffectedGroupsAndPeers(context.Background(), updated, 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(context.Background(), 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(context.Background(), 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"}, - PeerGroups: []string{"g2"}, - AccessControlGroups: []string{"g3"}, - } - groups, directPeers := collectRouteAffectedGroupsAndPeers(context.Background(), 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(context.Background(), r) - assert.ElementsMatch(t, []string{"g1"}, groups) - assert.ElementsMatch(t, []string{"p1"}, directPeers) -} - -func TestCollectRouteAffectedGroups_NilRoute(t *testing.T) { - groups, directPeers := collectRouteAffectedGroupsAndPeers(context.Background(), nil) - assert.Nil(t, groups) - assert.Nil(t, directPeers) -} - -func TestCollectRouteAffectedGroups_MultipleRoutes(t *testing.T) { - old := &route.Route{ - Groups: []string{"g1"}, - Peer: "p1", - } - updated := &route.Route{ - Groups: []string{"g2"}, - PeerGroups: []string{"g3"}, - } - groups, directPeers := collectRouteAffectedGroupsAndPeers(context.Background(), updated, 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) - }) - } -} +// Pure policy/route/router extraction unit tests moved to the affectedpeers +// package (management/server/affectedpeers) along with the logic they cover. // --------------------------------------------------------------------------- // resolvePeerIDs tests diff --git a/management/server/affectedpeers/resolver.go b/management/server/affectedpeers/resolver.go new file mode 100644 index 000000000..34f05ae80 --- /dev/null +++ b/management/server/affectedpeers/resolver.go @@ -0,0 +1,661 @@ +package affectedpeers + +import ( + "context" + + log "github.com/sirupsen/logrus" + + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" + 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" +) + +// Change describes what changed in an account. The resolver never consults the +// Enabled flag of any object: toggling Enabled is itself an observable change. +type Change struct { + ChangedGroupIDs []string + ChangedPeerIDs []string + Policies []*types.Policy + Routes []*route.Route + PostureCheckIDs []string + ResourceIDs []string + NetworkIDs []string +} + +func (c Change) isEmpty() bool { + return len(c.ChangedGroupIDs) == 0 && + len(c.ChangedPeerIDs) == 0 && + len(c.Policies) == 0 && + len(c.Routes) == 0 && + len(c.PostureCheckIDs) == 0 && + len(c.ResourceIDs) == 0 && + len(c.NetworkIDs) == 0 +} + +// Resolve returns the deduplicated peer IDs whose network map may have changed by +// the given Change. Safe to call inside or after a transaction. +// +// At trace level it logs the full reasoning — which inputs drove which graph +// walks to which groups/peers, including the resource<->router bridge hops — so a +// miscalculation can be diagnosed from the logs alone. +func Resolve(ctx context.Context, s store.Store, accountID string, c Change) ([]string, error) { + if c.isEmpty() { + return nil, nil + } + r := newResolver(ctx, s, accountID, c) + log.WithContext(ctx).Tracef("affectedpeers resolve start: account=%s changedGroups=%v changedPeers=%v policies=%d routes=%d postureChecks=%v resources=%v networks=%v", + accountID, c.ChangedGroupIDs, c.ChangedPeerIDs, len(c.Policies), len(c.Routes), c.PostureCheckIDs, c.ResourceIDs, c.NetworkIDs) + r.walk() + return r.expand() +} + +// Collect returns the affected group IDs and direct peer IDs without expanding +// groups to members. For tests asserting on the intermediate sets; use Resolve otherwise. +func Collect(ctx context.Context, s store.Store, accountID string, c Change) (groupIDs []string, directPeerIDs []string) { + if c.isEmpty() { + return nil, nil + } + r := newResolver(ctx, s, accountID, c) + r.walk() + return setToSlice(r.groupSet), setToSlice(r.peerSet) +} + +func newResolver(ctx context.Context, s store.Store, accountID string, c Change) *resolver { + r := &resolver{ + ctx: ctx, + store: s, + accountID: accountID, + change: c, + changedGroupSet: toSet(c.ChangedGroupIDs), + changedPeerSet: toSet(c.ChangedPeerIDs), + groupSet: make(map[string]struct{}), + peerSet: make(map[string]struct{}), + resourceIDs: toSet(c.ResourceIDs), + networkIDs: toSet(c.NetworkIDs), + } + r.matchedPolicies = append(r.matchedPolicies, c.Policies...) + return r +} + +func (r *resolver) walk() { + r.collectFromExplicitPolicies() + r.collectFromExplicitRoutes(r.change.Routes) + r.collectFromPostureChecks(r.change.PostureCheckIDs) + + if len(r.changedGroupSet) > 0 || len(r.changedPeerSet) > 0 { + r.collectFromPolicies() + r.collectFromRoutes() + r.collectFromNameServers() + r.collectFromDNSSettings() + r.collectFromNetworkRouters() + r.collectFromProxyServices() + } + + r.collectResourceRouterBridge() +} + +type resolver struct { + ctx context.Context + store store.Store + accountID string + change Change + + changedGroupSet map[string]struct{} + changedPeerSet map[string]struct{} + + groupSet map[string]struct{} + peerSet map[string]struct{} + + matchedPolicies []*types.Policy + resourceIDs map[string]struct{} + networkIDs map[string]struct{} + + // Memoized per-account collections: each is loaded from the store at most + // once per Resolve and only when a walker actually needs it. + cachedPolicies []*types.Policy + policiesLoaded bool + cachedResources []*resourceTypes.NetworkResource + resourcesLoaded bool + cachedRouters []*routerTypes.NetworkRouter + routersLoaded bool +} + +func (r *resolver) policies() []*types.Policy { + if r.policiesLoaded { + return r.cachedPolicies + } + r.policiesLoaded = true + policies, err := r.store.GetAccountPolicies(r.ctx, store.LockingStrengthNone, r.accountID) + if err != nil { + log.WithContext(r.ctx).Errorf("failed to get policies for affected peers resolution: %v", err) + return nil + } + r.cachedPolicies = policies + return r.cachedPolicies +} + +func (r *resolver) networkResources() []*resourceTypes.NetworkResource { + if r.resourcesLoaded { + return r.cachedResources + } + r.resourcesLoaded = true + resources, err := r.store.GetNetworkResourcesByAccountID(r.ctx, store.LockingStrengthNone, r.accountID) + if err != nil { + log.WithContext(r.ctx).Errorf("failed to get network resources for affected peers resolution: %v", err) + return nil + } + r.cachedResources = resources + return r.cachedResources +} + +func (r *resolver) networkRouters() []*routerTypes.NetworkRouter { + if r.routersLoaded { + return r.cachedRouters + } + r.routersLoaded = true + routers, err := r.store.GetNetworkRoutersByAccountID(r.ctx, store.LockingStrengthNone, r.accountID) + if err != nil { + log.WithContext(r.ctx).Errorf("failed to get network routers for affected peers resolution: %v", err) + return nil + } + r.cachedRouters = routers + return r.cachedRouters +} + +func (r *resolver) expand() ([]string, error) { + groupIDs := setToSlice(r.groupSet) + var peerIDs []string + if len(groupIDs) > 0 { + ids, err := r.store.GetPeerIDsByGroups(r.ctx, r.accountID, groupIDs) + if err != nil { + return nil, err + } + peerIDs = ids + } + + log.WithContext(r.ctx).Tracef("affectedpeers resolve expand: account=%s affectedGroups=%v -> %d group-member peers; direct peers=%v", + r.accountID, groupIDs, len(peerIDs), setToSlice(r.peerSet)) + + seen := make(map[string]struct{}, len(peerIDs)) + for _, id := range peerIDs { + seen[id] = struct{}{} + } + for id := range r.peerSet { + if _, ok := seen[id]; !ok { + peerIDs = append(peerIDs, id) + seen[id] = struct{}{} + } + } + + log.WithContext(r.ctx).Tracef("affectedpeers resolve done: account=%s -> %d affected peers: %v", r.accountID, len(peerIDs), peerIDs) + return peerIDs, nil +} + +func (r *resolver) collectFromExplicitPolicies() { + for _, policy := range r.matchedPolicies { + if policy == nil { + continue + } + log.WithContext(r.ctx).Tracef("collectFromExplicitPolicies: changed policy %s (%s) -> folding rule groups %v + direct peers", + policy.ID, policy.Name, policy.RuleGroups()) + addAll(r.groupSet, policy.RuleGroups()) + collectPolicyDirectPeers(policy, r.peerSet) + } +} + +func (r *resolver) collectFromExplicitRoutes(routes []*route.Route) { + for _, rt := range routes { + if rt == nil { + continue + } + log.WithContext(r.ctx).Tracef("collectFromExplicitRoutes: changed route %s -> folding groups=%v peerGroups=%v accessControlGroups=%v peer=%q", + rt.ID, rt.Groups, rt.PeerGroups, rt.AccessControlGroups, rt.Peer) + addAll(r.groupSet, rt.Groups, rt.PeerGroups, rt.AccessControlGroups) + if rt.Peer != "" { + r.peerSet[rt.Peer] = struct{}{} + } + } +} + +func (r *resolver) collectFromPostureChecks(postureCheckIDs []string) { + if len(postureCheckIDs) == 0 { + return + } + ids := toSet(postureCheckIDs) + for _, policy := range r.policies() { + if !policyReferencesPostureChecks(policy, ids) { + continue + } + log.WithContext(r.ctx).Tracef("collectFromPostureChecks: policy %s (%s) references changed posture checks %v -> folding rule groups %v + direct peers", + policy.ID, policy.Name, postureCheckIDs, policy.RuleGroups()) + addAll(r.groupSet, policy.RuleGroups()) + collectPolicyDirectPeers(policy, r.peerSet) + r.matchedPolicies = append(r.matchedPolicies, policy) + } +} + +func (r *resolver) collectFromPolicies() { + for _, policy := range r.policies() { + matchedByGroup := policyReferencesGroups(policy, r.changedGroupSet) + matchedByPeer := len(r.changedPeerSet) > 0 && policyReferencesDirectPeers(policy, r.changedPeerSet) + if !matchedByGroup && !matchedByPeer { + continue + } + log.WithContext(r.ctx).Tracef("collectFromPolicies: policy %s (%s) matched (byGroup=%t byPeer=%t) -> folding rule groups %v + direct peers", + policy.ID, policy.Name, matchedByGroup, matchedByPeer, policy.RuleGroups()) + addAll(r.groupSet, policy.RuleGroups()) + collectPolicyDirectPeers(policy, r.peerSet) + r.matchedPolicies = append(r.matchedPolicies, policy) + } +} + +func (r *resolver) collectFromRoutes() { + routes, err := r.store.GetAccountRoutes(r.ctx, store.LockingStrengthNone, r.accountID) + if err != nil { + log.WithContext(r.ctx).Errorf("failed to get routes for affected peers resolution: %v", err) + return + } + for _, rt := range routes { + matchedByGroup := anyInSet(rt.Groups, r.changedGroupSet) || anyInSet(rt.PeerGroups, r.changedGroupSet) || anyInSet(rt.AccessControlGroups, r.changedGroupSet) + matchedByPeer := rt.Peer != "" && len(r.changedPeerSet) > 0 && isInSet(rt.Peer, r.changedPeerSet) + if !matchedByGroup && !matchedByPeer { + continue + } + log.WithContext(r.ctx).Tracef("collectFromRoutes: route %s matched (byGroup=%t byPeer=%t) -> folding groups=%v peerGroups=%v accessControlGroups=%v peer=%q", + rt.ID, matchedByGroup, matchedByPeer, rt.Groups, rt.PeerGroups, rt.AccessControlGroups, rt.Peer) + addAll(r.groupSet, rt.Groups, rt.PeerGroups, rt.AccessControlGroups) + if rt.Peer != "" { + r.peerSet[rt.Peer] = struct{}{} + } + } +} + +func (r *resolver) collectFromNameServers() { + if len(r.changedGroupSet) == 0 { + return + } + nsGroups, err := r.store.GetAccountNameServerGroups(r.ctx, store.LockingStrengthNone, r.accountID) + if err != nil { + log.WithContext(r.ctx).Errorf("failed to get nameserver groups for affected peers resolution: %v", err) + return + } + for _, ns := range nsGroups { + if anyInSet(ns.Groups, r.changedGroupSet) { + log.WithContext(r.ctx).Tracef("collectFromNameServers: nameserver group %s references a changed group -> folding its groups %v", ns.ID, ns.Groups) + addAll(r.groupSet, ns.Groups) + } + } +} + +func (r *resolver) collectFromDNSSettings() { + if len(r.changedGroupSet) == 0 { + return + } + dnsSettings, err := r.store.GetAccountDNSSettings(r.ctx, store.LockingStrengthNone, r.accountID) + if err != nil { + log.WithContext(r.ctx).Errorf("failed to get DNS settings for affected peers resolution: %v", err) + return + } + for _, gID := range dnsSettings.DisabledManagementGroups { + if _, ok := r.changedGroupSet[gID]; ok { + log.WithContext(r.ctx).Tracef("collectFromDNSSettings: changed group %s is in DisabledManagementGroups -> folding it", gID) + r.groupSet[gID] = struct{}{} + } + } +} + +func (r *resolver) collectFromNetworkRouters() { + for _, router := range r.networkRouters() { + matchedByGroup := anyInSet(router.PeerGroups, r.changedGroupSet) + matchedByPeer := router.Peer != "" && len(r.changedPeerSet) > 0 && isInSet(router.Peer, r.changedPeerSet) + if !matchedByGroup && !matchedByPeer { + continue + } + log.WithContext(r.ctx).Tracef("collectFromNetworkRouters: router %s on network %s matched (byGroup=%t byPeer=%t) -> folding peerGroups=%v peer=%q and marking network for source bridge", + router.ID, router.NetworkID, matchedByGroup, matchedByPeer, router.PeerGroups, router.Peer) + addAll(r.groupSet, router.PeerGroups) + if router.Peer != "" { + r.peerSet[router.Peer] = struct{}{} + } + r.networkIDs[router.NetworkID] = struct{}{} + } +} + +func (r *resolver) collectFromProxyServices() { + services, proxyByCluster, ok := r.loadProxyServiceContext() + if !ok { + return + } + + expanded := r.expandChangedPeersWithGroups() + + for _, svc := range services { + if svc == nil { + continue + } + proxyPeers := proxyByCluster[svc.ProxyCluster] + if len(proxyPeers) == 0 { + continue + } + matchedByPeer := serviceMatchesChangedPeers(svc, proxyPeers, expanded) + matchedByAccessGroup := anyInSet(svc.AccessGroups, r.changedGroupSet) + if !matchedByPeer && !matchedByAccessGroup { + continue + } + log.WithContext(r.ctx).Tracef("collectFromProxyServices: service %s (cluster=%s) matched (byProxyOrTargetPeer=%t byAccessGroup=%t) -> folding %d proxy peers, peer targets and access groups %v", + svc.ID, svc.ProxyCluster, matchedByPeer, matchedByAccessGroup, len(proxyPeers), svc.AccessGroups) + for _, pid := range proxyPeers { + r.peerSet[pid] = struct{}{} + } + for _, target := range svc.Targets { + if target.TargetType == rpservice.TargetTypePeer && target.TargetId != "" { + r.peerSet[target.TargetId] = struct{}{} + } + } + addAll(r.groupSet, svc.AccessGroups) + } +} + +func (r *resolver) loadProxyServiceContext() ([]*rpservice.Service, map[string][]string, bool) { + // Embedded proxy peers are the prerequisite for any synthesized proxy policy. + // Probe that first (a narrow, indexed lookup) and skip the services table load + // entirely when the account has no embedded proxy peers. + proxyByCluster, err := r.store.GetEmbeddedProxyPeerIDsByCluster(r.ctx, r.accountID) + if err != nil { + log.WithContext(r.ctx).Errorf("failed to get embedded proxy peers for affected peers resolution: %v", err) + return nil, nil, false + } + if len(proxyByCluster) == 0 { + return nil, nil, false + } + services, err := r.store.GetAccountServices(r.ctx, store.LockingStrengthNone, r.accountID) + if err != nil { + log.WithContext(r.ctx).Errorf("failed to get services for affected peers resolution: %v", err) + return nil, nil, false + } + if len(services) == 0 { + return nil, nil, false + } + return services, proxyByCluster, true +} + +func (r *resolver) expandChangedPeersWithGroups() map[string]struct{} { + if len(r.changedGroupSet) == 0 { + return r.changedPeerSet + } + ids, err := r.store.GetPeerIDsByGroups(r.ctx, r.accountID, setToSlice(r.changedGroupSet)) + if err != nil { + log.WithContext(r.ctx).Errorf("failed to expand changed groups to peers for service resolution: %v", err) + return r.changedPeerSet + } + if len(ids) == 0 { + return r.changedPeerSet + } + merged := make(map[string]struct{}, len(r.changedPeerSet)+len(ids)) + for id := range r.changedPeerSet { + merged[id] = struct{}{} + } + for _, id := range ids { + merged[id] = struct{}{} + } + return merged +} + +// collectResourceRouterBridge folds in the routing peers serving the resources +// targeted by matched/explicit policies (source -> router), and the source peers +// of policies serving resources on the affected networks (router -> source). The +// routing peer is reachable only through resource -> network -> router, never +// through the policy's own groups, so it must be collected here. +func (r *resolver) collectResourceRouterBridge() { + r.bridgeSourceToRouters() + r.bridgeRoutersToSources() +} + +func (r *resolver) bridgeSourceToRouters() { + resourceIDs := policyDestinationResourceIDs(r.ctx, r.store, r.accountID, r.matchedPolicies...) + for id := range r.resourceIDs { + resourceIDs[id] = struct{}{} + } + if len(resourceIDs) == 0 { + return + } + + networkIDs := r.resourceNetworkIDs(resourceIDs) + log.WithContext(r.ctx).Tracef("bridgeSourceToRouters: targeted resources %v -> networks %v (their routers become affected via the router->source pass)", + setToSlice(resourceIDs), setToSlice(networkIDs)) + for id := range networkIDs { + r.networkIDs[id] = struct{}{} + } +} + +func (r *resolver) bridgeRoutersToSources() { + if len(r.networkIDs) == 0 { + return + } + + log.WithContext(r.ctx).Tracef("bridgeRoutersToSources: affected networks %v -> folding their routing peers and the source peers of policies targeting their resources", + setToSlice(r.networkIDs)) + + r.foldRoutersOnNetworks(r.networkIDs) + + resourceIDs := make(map[string]struct{}) + for _, resource := range r.networkResources() { + if _, ok := r.networkIDs[resource.NetworkID]; ok { + resourceIDs[resource.ID] = struct{}{} + } + } + if len(resourceIDs) == 0 { + return + } + + for _, policy := range r.policies() { + if r.policyTargetsResources(policy, resourceIDs) { + log.WithContext(r.ctx).Tracef("bridgeRoutersToSources: policy %s (%s) targets an affected-network resource -> folding its source groups/peers", policy.ID, policy.Name) + collectPolicySources(policy, r.groupSet, r.peerSet) + } + } +} + +func (r *resolver) foldRoutersOnNetworks(networkIDs map[string]struct{}) { + for _, router := range r.networkRouters() { + if _, ok := networkIDs[router.NetworkID]; !ok { + continue + } + log.WithContext(r.ctx).Tracef("bridgeRoutersToSources: router %s serves affected network %s -> folding peerGroups=%v peer=%q", + router.ID, router.NetworkID, router.PeerGroups, router.Peer) + addAll(r.groupSet, router.PeerGroups) + if router.Peer != "" { + r.peerSet[router.Peer] = struct{}{} + } + } +} + +func (r *resolver) resourceNetworkIDs(resourceIDs map[string]struct{}) map[string]struct{} { + networkIDs := make(map[string]struct{}) + for _, resource := range r.networkResources() { + if _, ok := resourceIDs[resource.ID]; ok { + networkIDs[resource.NetworkID] = struct{}{} + } + } + return networkIDs +} + +func (r *resolver) policyTargetsResources(policy *types.Policy, resourceIDs map[string]struct{}) bool { + if policy == nil { + return false + } + destGroupSet := make(map[string]struct{}) + for _, rule := range policy.Rules { + if rule.DestinationResource.Type != types.ResourceTypePeer && isInSet(rule.DestinationResource.ID, resourceIDs) { + return true + } + for _, gID := range rule.Destinations { + destGroupSet[gID] = struct{}{} + } + } + if len(destGroupSet) == 0 { + return false + } + groups, err := r.store.GetGroupsByIDs(r.ctx, store.LockingStrengthNone, r.accountID, setToSlice(destGroupSet)) + if err != nil { + log.WithContext(r.ctx).Errorf("failed to get destination groups for router policy bridge: %v", err) + return false + } + for _, group := range groups { + for _, res := range group.Resources { + if isInSet(res.ID, resourceIDs) { + return true + } + } + } + return false +} + +func policyDestinationResourceIDs(ctx context.Context, s store.Store, accountID string, policies ...*types.Policy) map[string]struct{} { + destGroupSet := make(map[string]struct{}) + resourceIDs := make(map[string]struct{}) + + for _, policy := range policies { + if policy == nil { + continue + } + for _, rule := range policy.Rules { + for _, gID := range rule.Destinations { + destGroupSet[gID] = struct{}{} + } + if rule.DestinationResource.Type != types.ResourceTypePeer && rule.DestinationResource.ID != "" { + resourceIDs[rule.DestinationResource.ID] = struct{}{} + } + } + } + + if len(destGroupSet) > 0 { + groups, err := s.GetGroupsByIDs(ctx, store.LockingStrengthNone, accountID, setToSlice(destGroupSet)) + if err != nil { + log.WithContext(ctx).Errorf("failed to get destination groups for resource router bridge: %v", err) + } else { + for _, group := range groups { + for _, res := range group.Resources { + if res.ID != "" { + resourceIDs[res.ID] = struct{}{} + } + } + } + } + } + + return resourceIDs +} + +func collectPolicyDirectPeers(policy *types.Policy, peerSet map[string]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{}{} + } + } +} + +func collectPolicySources(policy *types.Policy, groupSet, peerSet map[string]struct{}) { + for _, rule := range policy.Rules { + addAll(groupSet, rule.Sources) + if rule.SourceResource.Type == types.ResourceTypePeer && rule.SourceResource.ID != "" { + peerSet[rule.SourceResource.ID] = 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 policyReferencesDirectPeers(policy *types.Policy, changedSet map[string]struct{}) bool { + for _, rule := range policy.Rules { + if isDirectPeerInSet(rule.SourceResource, changedSet) || isDirectPeerInSet(rule.DestinationResource, changedSet) { + return true + } + } + return false +} + +func policyReferencesPostureChecks(policy *types.Policy, ids map[string]struct{}) bool { + for _, id := range policy.SourcePostureChecks { + if _, ok := ids[id]; ok { + 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 serviceMatchesChangedPeers(svc *rpservice.Service, proxyPeers []string, changedPeers map[string]struct{}) bool { + for _, pid := range proxyPeers { + if _, ok := changedPeers[pid]; ok { + return true + } + } + for _, target := range svc.Targets { + if target.TargetType != rpservice.TargetTypePeer || target.TargetId == "" { + continue + } + if _, ok := changedPeers[target.TargetId]; ok { + return true + } + } + return false +} + +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 addAll(set map[string]struct{}, slices ...[]string) { + for _, s := range slices { + for _, id := range s { + set[id] = struct{}{} + } + } +} + +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/affectedpeers/resolver_test.go b/management/server/affectedpeers/resolver_test.go new file mode 100644 index 000000000..2f9c2d523 --- /dev/null +++ b/management/server/affectedpeers/resolver_test.go @@ -0,0 +1,138 @@ +package affectedpeers + +import ( + "testing" + + "github.com/stretchr/testify/assert" + + "github.com/netbirdio/netbird/management/server/types" +) + +// policyGroupsAndPeers mirrors the explicit-policy extraction (RuleGroups + +// direct peers) the resolver folds in, for asserting the pure logic. +func policyGroupsAndPeers(policies ...*types.Policy) (groups []string, peers []string) { + peerSet := map[string]struct{}{} + for _, p := range policies { + if p == nil { + continue + } + groups = append(groups, p.RuleGroups()...) + collectPolicyDirectPeers(p, peerSet) + } + for id := range peerSet { + peers = append(peers, id) + } + return groups, peers +} + +func TestPolicyGroupsAndPeers_Basic(t *testing.T) { + policy := &types.Policy{Rules: []*types.PolicyRule{{Sources: []string{"g1", "g2"}, Destinations: []string{"g3"}}}} + groups, peers := policyGroupsAndPeers(policy) + assert.ElementsMatch(t, []string{"g1", "g2", "g3"}, groups) + assert.Empty(t, peers) +} + +func TestPolicyGroupsAndPeers_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, peers := policyGroupsAndPeers(policy) + assert.ElementsMatch(t, []string{"g1", "g2"}, groups) + assert.ElementsMatch(t, []string{"p1", "p2"}, peers) +} + +func TestPolicyGroupsAndPeers_NilPolicy(t *testing.T) { + groups, peers := policyGroupsAndPeers(nil) + assert.Nil(t, groups) + assert.Nil(t, peers) +} + +func TestPolicyGroupsAndPeers_MultiplePolicies(t *testing.T) { + old := &types.Policy{Rules: []*types.PolicyRule{{Sources: []string{"g1"}, Destinations: []string{"g2"}}}} + updated := &types.Policy{Rules: []*types.PolicyRule{{Sources: []string{"g3"}, Destinations: []string{"g4"}}}} + groups, _ := policyGroupsAndPeers(updated, old) + assert.ElementsMatch(t, []string{"g1", "g2", "g3", "g4"}, groups) +} + +func TestPolicyGroupsAndPeers_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, peers := policyGroupsAndPeers(policy) + assert.ElementsMatch(t, []string{"g1", "g2"}, groups) + assert.Empty(t, peers, "domain resource type should not produce direct peer IDs") +} + +func TestChangeIsEmpty(t *testing.T) { + assert.True(t, Change{}.isEmpty()) + assert.False(t, Change{ChangedGroupIDs: []string{"g"}}.isEmpty()) + assert.False(t, Change{ChangedPeerIDs: []string{"p"}}.isEmpty()) + assert.False(t, Change{Policies: []*types.Policy{{}}}.isEmpty()) + assert.False(t, Change{ResourceIDs: []string{"r"}}.isEmpty()) + assert.False(t, Change{NetworkIDs: []string{"n"}}.isEmpty()) + assert.False(t, Change{PostureCheckIDs: []string{"pc"}}.isEmpty()) +} + +func TestPolicyReferencesGroups(t *testing.T) { + policy := &types.Policy{Rules: []*types.PolicyRule{{Sources: []string{"g1", "g2"}, Destinations: []string{"g3"}}}} + + assert.True(t, policyReferencesGroups(policy, map[string]struct{}{"g1": {}})) + assert.True(t, policyReferencesGroups(policy, map[string]struct{}{"g3": {}})) + assert.False(t, policyReferencesGroups(policy, map[string]struct{}{"g4": {}})) + assert.False(t, policyReferencesGroups(policy, map[string]struct{}{})) +} + +func TestPolicyReferencesDirectPeers(t *testing.T) { + policy := &types.Policy{Rules: []*types.PolicyRule{{ + SourceResource: types.Resource{Type: types.ResourceTypePeer, ID: "p1"}, + DestinationResource: types.Resource{Type: types.ResourceTypeHost, ID: "r1"}, + }}} + + assert.True(t, policyReferencesDirectPeers(policy, map[string]struct{}{"p1": {}})) + assert.False(t, policyReferencesDirectPeers(policy, map[string]struct{}{"r1": {}})) + assert.False(t, policyReferencesDirectPeers(policy, map[string]struct{}{"p2": {}})) +} + +func TestPolicyReferencesPostureChecks(t *testing.T) { + policy := &types.Policy{SourcePostureChecks: []string{"pc1", "pc2"}} + + assert.True(t, policyReferencesPostureChecks(policy, map[string]struct{}{"pc1": {}})) + assert.False(t, policyReferencesPostureChecks(policy, map[string]struct{}{"pc3": {}})) +} + +func TestCollectPolicyDirectPeers(t *testing.T) { + policy := &types.Policy{Rules: []*types.PolicyRule{{ + SourceResource: types.Resource{Type: types.ResourceTypePeer, ID: "p1"}, + DestinationResource: types.Resource{Type: types.ResourceTypePeer, ID: "p2"}, + }, { + DestinationResource: types.Resource{Type: types.ResourceTypeHost, ID: "r1"}, + }}} + + peerSet := map[string]struct{}{} + collectPolicyDirectPeers(policy, peerSet) + + assert.Contains(t, peerSet, "p1") + assert.Contains(t, peerSet, "p2") + assert.NotContains(t, peerSet, "r1") +} + +func TestCollectPolicySources(t *testing.T) { + policy := &types.Policy{Rules: []*types.PolicyRule{{ + Sources: []string{"g1"}, + SourceResource: types.Resource{Type: types.ResourceTypePeer, ID: "p1"}, + Destinations: []string{"g2"}, + }}} + + groupSet := map[string]struct{}{} + peerSet := map[string]struct{}{} + collectPolicySources(policy, groupSet, peerSet) + + assert.Contains(t, groupSet, "g1") + assert.NotContains(t, groupSet, "g2", "destination groups must not be collected as sources") + assert.Contains(t, peerSet, "p1") +} diff --git a/management/server/group.go b/management/server/group.go index f1288000e..7c3da45b1 100644 --- a/management/server/group.go +++ b/management/server/group.go @@ -10,6 +10,7 @@ import ( log "github.com/sirupsen/logrus" "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/affectedpeers" "github.com/netbirdio/netbird/management/server/permissions/modules" "github.com/netbirdio/netbird/management/server/permissions/operations" "github.com/netbirdio/netbird/management/server/store" @@ -98,8 +99,7 @@ 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) + affectedPeerIDs = am.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{ChangedGroupIDs: []string{newGroup.ID}}) return transaction.IncrementNetworkSerial(ctx, accountID) }) @@ -171,8 +171,7 @@ func (am *DefaultAccountManager) UpdateGroup(ctx context.Context, accountID, use return err } - groupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, transaction, accountID, []string{newGroup.ID}) - affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, append(directPeerIDs, peersToRemove...)) + affectedPeerIDs = am.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{ChangedGroupIDs: []string{newGroup.ID}, ChangedPeerIDs: peersToRemove}) return transaction.IncrementNetworkSerial(ctx, accountID) }) @@ -249,8 +248,7 @@ func (am *DefaultAccountManager) CreateGroups(ctx context.Context, accountID, us storeEvent() } - allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, am.Store, accountID, groupIDs) - affectedPeerIDs := am.resolvePeerIDs(ctx, am.Store, accountID, allGroupIDs, directPeerIDs) + affectedPeerIDs := am.ResolveAffectedPeers(ctx, am.Store, accountID, affectedpeers.Change{ChangedGroupIDs: groupIDs}) if len(affectedPeerIDs) > 0 { log.WithContext(ctx).Debugf("CreateGroups %v: updating %d affected peers: %v", groupIDs, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) @@ -296,8 +294,7 @@ func (am *DefaultAccountManager) UpdateGroups(ctx context.Context, accountID, us storeEvent() } - allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, am.Store, accountID, groupIDs) - affectedPeerIDs := am.resolvePeerIDs(ctx, am.Store, accountID, allGroupIDs, directPeerIDs) + affectedPeerIDs := am.ResolveAffectedPeers(ctx, am.Store, accountID, affectedpeers.Change{ChangedGroupIDs: groupIDs}) if len(affectedPeerIDs) > 0 { log.WithContext(ctx).Debugf("UpdateGroups %v: updating %d affected peers: %v", groupIDs, len(affectedPeerIDs), affectedPeerIDs) am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) @@ -435,6 +432,7 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us var allErrors error var groupIDsToDelete []string var deletedGroups []*types.Group + var affectedPeerIDs []string extraSettings, err := am.settingsManager.GetExtraSettings(ctx, accountID) if err != nil { @@ -462,6 +460,8 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us return allErrors } + affectedPeerIDs = am.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{ChangedGroupIDs: groupIDsToDelete}) + if err = transaction.DeleteGroups(ctx, accountID, groupIDsToDelete); err != nil { return err } @@ -480,6 +480,13 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us am.StoreEvent(ctx, userID, group.ID, accountID, activity.GroupDeleted, group.EventMeta()) } + if len(affectedPeerIDs) > 0 { + log.WithContext(ctx).Debugf("DeleteGroups %v: updating %d affected peers: %v", groupIDsToDelete, len(affectedPeerIDs), affectedPeerIDs) + am.UpdateAffectedPeers(ctx, accountID, affectedPeerIDs) + } else { + log.WithContext(ctx).Tracef("DeleteGroups %v: no affected peers", groupIDsToDelete) + } + return allErrors } @@ -497,8 +504,7 @@ func (am *DefaultAccountManager) GroupAddPeer(ctx context.Context, accountID, gr return err } - allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, transaction, accountID, []string{groupID}) - affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, allGroupIDs, directPeerIDs) + affectedPeerIDs = am.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{ChangedGroupIDs: []string{groupID}}) return transaction.IncrementNetworkSerial(ctx, accountID) }) @@ -536,8 +542,7 @@ func (am *DefaultAccountManager) GroupAddResource(ctx context.Context, accountID return err } - allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, transaction, accountID, []string{groupID}) - affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, allGroupIDs, directPeerIDs) + affectedPeerIDs = am.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{ChangedGroupIDs: []string{groupID}}) return transaction.IncrementNetworkSerial(ctx, accountID) }) @@ -562,8 +567,7 @@ func (am *DefaultAccountManager) GroupDeletePeer(ctx context.Context, accountID, err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { // 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) + affectedPeerIDs = am.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{ChangedGroupIDs: []string{groupID}}) if err = transaction.RemovePeerFromGroup(ctx, peerID, groupID); err != nil { return err @@ -609,8 +613,7 @@ func (am *DefaultAccountManager) GroupDeleteResource(ctx context.Context, accoun return err } - allGroupIDs, directPeerIDs := collectGroupChangeAffectedGroups(ctx, transaction, accountID, []string{groupID}) - affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, allGroupIDs, directPeerIDs) + affectedPeerIDs = am.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{ChangedGroupIDs: []string{groupID}}) return transaction.IncrementNetworkSerial(ctx, accountID) }) diff --git a/management/server/group_linkage.go b/management/server/group_linkage.go index 626ef7956..be1d56c2d 100644 --- a/management/server/group_linkage.go +++ b/management/server/group_linkage.go @@ -285,3 +285,12 @@ func networkRoutersReferenceGroups(ctx context.Context, transaction store.Store, } 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 +} diff --git a/management/server/mock_server/account_mock.go b/management/server/mock_server/account_mock.go index 28132a75b..894f7ce1e 100644 --- a/management/server/mock_server/account_mock.go +++ b/management/server/mock_server/account_mock.go @@ -15,6 +15,7 @@ import ( "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" "github.com/netbirdio/netbird/management/server/account" "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/affectedpeers" "github.com/netbirdio/netbird/management/server/idp" nbpeer "github.com/netbirdio/netbird/management/server/peer" "github.com/netbirdio/netbird/management/server/posture" @@ -134,6 +135,7 @@ type MockAccountManager struct { 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, reason types.UpdateReason) + ResolveAffectedPeersFunc func(ctx context.Context, s store.Store, accountID string, change affectedpeers.Change) []string BufferUpdateAccountPeersFunc func(ctx context.Context, accountID string, reason types.UpdateReason) RecalculateNetworkMapCacheFunc func(ctx context.Context, accountId string) error @@ -217,6 +219,13 @@ func (am *MockAccountManager) UpdateAffectedPeers(ctx context.Context, accountID } } +func (am *MockAccountManager) ResolveAffectedPeers(ctx context.Context, s store.Store, accountID string, change affectedpeers.Change) []string { + if am.ResolveAffectedPeersFunc != nil { + return am.ResolveAffectedPeersFunc(ctx, s, accountID, change) + } + return nil +} + func (am *MockAccountManager) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string, reason types.UpdateReason) { if am.BufferUpdateAffectedPeersFunc != nil { am.BufferUpdateAffectedPeersFunc(ctx, accountID, peerIDs, reason) diff --git a/management/server/networks/manager.go b/management/server/networks/manager.go index e485ec15a..d96b5d447 100644 --- a/management/server/networks/manager.go +++ b/management/server/networks/manager.go @@ -9,6 +9,7 @@ import ( "github.com/netbirdio/netbird/management/server/account" "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/affectedpeers" "github.com/netbirdio/netbird/management/server/networks/resources" "github.com/netbirdio/netbird/management/server/networks/routers" "github.com/netbirdio/netbird/management/server/networks/types" @@ -16,7 +17,6 @@ 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" ) @@ -113,14 +113,6 @@ 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, ctx, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.Networks, operations.Delete) if err != nil { @@ -136,22 +128,16 @@ func (m *managerImpl) DeleteNetwork(ctx context.Context, accountID, userID, netw } var eventsToStore []func() - var affectedData *networkAffectedPeersData + var affectedPeerIDs []string err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { + affectedPeerIDs = m.accountManager.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{NetworkIDs: []string{networkID}}) + 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) @@ -164,14 +150,7 @@ func (m *managerImpl) DeleteNetwork(ctx context.Context, accountID, userID, netw return fmt.Errorf("failed to get routers in network: %w", err) } - 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) @@ -179,24 +158,6 @@ 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) @@ -221,111 +182,16 @@ func (m *managerImpl) DeleteNetwork(ctx context.Context, accountID, userID, netw event() } - 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) - } + 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) } 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 { - 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 { - groupSet[gID] = struct{}{} - } - - if len(data.resourceGroupIDs) > 0 { - for _, gID := range data.resourceGroupIDs { - groupSet[gID] = struct{}{} - } - 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) - } - - 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 { - 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 NewManagerMock() Manager { return &mockManager{} } diff --git a/management/server/networks/resources/manager.go b/management/server/networks/resources/manager.go index ea0f86727..03b37db6f 100644 --- a/management/server/networks/resources/manager.go +++ b/management/server/networks/resources/manager.go @@ -10,6 +10,7 @@ import ( "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" "github.com/netbirdio/netbird/management/server/account" "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/affectedpeers" "github.com/netbirdio/netbird/management/server/groups" "github.com/netbirdio/netbird/management/server/networks/resources/types" "github.com/netbirdio/netbird/management/server/permissions" @@ -114,10 +115,10 @@ func (m *managerImpl) CreateResource(ctx context.Context, userID string, resourc } var eventsToStore []func() - var affectedData *resourceAffectedPeersData + var affectedPeerIDs []string err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { var txErr error - eventsToStore, affectedData, txErr = m.createResourceInTransaction(ctx, transaction, userID, resource) + eventsToStore, affectedPeerIDs, txErr = m.createResourceInTransaction(ctx, transaction, userID, resource) return txErr }) if err != nil { @@ -128,7 +129,7 @@ func (m *managerImpl) CreateResource(ctx context.Context, userID string, resourc event() } - if affectedPeerIDs := m.resolveResourceAffectedPeers(ctx, resource.AccountID, affectedData); len(affectedPeerIDs) > 0 { + if 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 { @@ -138,7 +139,7 @@ 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) { +func (m *managerImpl) createResourceInTransaction(ctx context.Context, transaction store.Store, userID string, resource *types.NetworkResource) ([]func(), []string, 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) @@ -174,12 +175,9 @@ func (m *managerImpl) createResourceInTransaction(ctx context.Context, transacti 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) - } + affectedPeerIDs := m.accountManager.ResolveAffectedPeers(ctx, transaction, resource.AccountID, affectedpeers.Change{ResourceIDs: []string{resource.ID}}) - return eventsToStore, affectedData, nil + return eventsToStore, affectedPeerIDs, nil } func (m *managerImpl) GetResource(ctx context.Context, accountID, userID, networkID, resourceID string) (*types.NetworkResource, error) { @@ -222,7 +220,7 @@ func (m *managerImpl) UpdateResource(ctx context.Context, userID string, resourc resource.Prefix = prefix var eventsToStore []func() - var affectedData *resourceAffectedPeersData + var affectedPeerIDs []string err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { network, err := transaction.GetNetworkByID(ctx, store.LockingStrengthUpdate, resource.AccountID, resource.NetworkID) if err != nil { @@ -248,7 +246,7 @@ 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) + oldGroups, err := m.groupsManager.GetResourceGroupsInTransaction(ctx, transaction, store.LockingStrengthNone, resource.AccountID, resource.ID) if err != nil { return fmt.Errorf("failed to get old resource groups: %w", err) } @@ -272,10 +270,12 @@ 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) - } + // Pass both old and new resource group IDs so policies that targeted the + // resource via a now-detached group still refresh their source peers. + affectedPeerIDs = m.accountManager.ResolveAffectedPeers(ctx, transaction, resource.AccountID, affectedpeers.Change{ + ResourceIDs: []string{resource.ID}, + ChangedGroupIDs: append(oldGroupIDs, resource.GroupIDs...), + }) err = transaction.IncrementNetworkSerial(ctx, resource.AccountID) if err != nil { @@ -300,7 +300,7 @@ func (m *managerImpl) UpdateResource(ctx context.Context, userID string, resourc } }() - if affectedPeerIDs := m.resolveResourceAffectedPeers(ctx, resource.AccountID, affectedData); len(affectedPeerIDs) > 0 { + if 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 { @@ -366,21 +366,9 @@ func (m *managerImpl) DeleteResource(ctx context.Context, accountID, userID, net } var events []func() - var affectedData *resourceAffectedPeersData + var affectedPeerIDs []string 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) - } + affectedPeerIDs = m.accountManager.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{ResourceIDs: []string{resourceID}}) events, err = m.DeleteResourceInTransaction(ctx, transaction, accountID, userID, networkID, resourceID) if err != nil { @@ -402,7 +390,7 @@ func (m *managerImpl) DeleteResource(ctx context.Context, accountID, userID, net event() } - if affectedPeerIDs := m.resolveResourceAffectedPeers(ctx, accountID, affectedData); len(affectedPeerIDs) > 0 { + if 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 { @@ -454,151 +442,6 @@ 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 &resourceAffectedPeersData{}, 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 - } - - 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{}) - directPeerIDs := collectResourcePolicySourceGroups(data.policies, data.resourceGroupIDs, groupSet) - - for _, gID := range data.routerPeerGroups { - groupSet[gID] = struct{}{} - } - directPeerIDs = append(directPeerIDs, data.routerDirectPeers...) - - if len(groupSet) == 0 && len(directPeerIDs) == 0 { - 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 - } - 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 -} - -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 := 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 { - 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 NewManagerMock() Manager { return &mockManager{} } diff --git a/management/server/networks/routers/manager.go b/management/server/networks/routers/manager.go index f5bb732bc..7384cab5f 100644 --- a/management/server/networks/routers/manager.go +++ b/management/server/networks/routers/manager.go @@ -10,13 +10,13 @@ import ( "github.com/netbirdio/netbird/management/server/account" "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/affectedpeers" "github.com/netbirdio/netbird/management/server/networks/routers/types" networkTypes "github.com/netbirdio/netbird/management/server/networks/types" "github.com/netbirdio/netbird/management/server/permissions" "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" ) @@ -91,7 +91,7 @@ func (m *managerImpl) CreateRouter(ctx context.Context, userID string, router *t } var network *networkTypes.Network - var affectedData *routerAffectedPeersData + var affectedPeerIDs []string err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { network, err = transaction.GetNetworkByID(ctx, store.LockingStrengthNone, router.AccountID, router.NetworkID) if err != nil { @@ -114,10 +114,7 @@ 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) - } + affectedPeerIDs = m.accountManager.ResolveAffectedPeers(ctx, transaction, router.AccountID, affectedpeers.Change{NetworkIDs: []string{router.NetworkID}}) return nil }) @@ -127,7 +124,7 @@ 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 { + if 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 { @@ -168,10 +165,10 @@ func (m *managerImpl) UpdateRouter(ctx context.Context, userID string, router *t } var network *networkTypes.Network - var affectedData *routerAffectedPeersData + var affectedPeerIDs []string err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { var txErr error - network, affectedData, txErr = m.updateRouterInTransaction(ctx, transaction, router) + network, affectedPeerIDs, txErr = m.updateRouterInTransaction(ctx, transaction, router) return txErr }) if err != nil { @@ -180,7 +177,7 @@ 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 { + if 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 { @@ -190,7 +187,7 @@ 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) { +func (m *managerImpl) updateRouterInTransaction(ctx context.Context, transaction store.Store, router *types.NetworkRouter) (*networkTypes.Network, []string, 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) @@ -209,16 +206,6 @@ func (m *managerImpl) updateRouterInTransaction(ctx context.Context, transaction return nil, nil, status.NewRouterNotPartOfNetworkError(router.ID, router.NetworkID) } - allPeerGroups := append([]string{}, router.PeerGroups...) - allPeerGroups = append(allPeerGroups, existing.PeerGroups...) - var directPeers []string - if router.Peer != "" { - directPeers = append(directPeers, router.Peer) - } - if existing.Peer != "" { - directPeers = append(directPeers, existing.Peer) - } - if err = transaction.UpdateNetworkRouter(ctx, router); err != nil { return nil, nil, fmt.Errorf("failed to update network router: %w", err) } @@ -227,12 +214,37 @@ func (m *managerImpl) updateRouterInTransaction(ctx context.Context, transaction 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) + networkIDs := []string{router.NetworkID} + if existing.NetworkID != router.NetworkID { + networkIDs = append(networkIDs, existing.NetworkID) } - return network, affectedData, nil + affectedPeerIDs := m.accountManager.ResolveAffectedPeers(ctx, transaction, router.AccountID, affectedpeers.Change{NetworkIDs: networkIDs}) + + // The previous routing peer / peer-group members lose their routing role and + // are no longer reachable from the post-update network state, so add them + // explicitly. + affectedPeerIDs = append(affectedPeerIDs, oldRoutingPeerIDs(ctx, transaction, router.AccountID, existing)...) + + return network, affectedPeerIDs, nil +} + +// oldRoutingPeerIDs returns the peer IDs that served as the router's routing peers +// before an update (direct Peer plus PeerGroups members). +func oldRoutingPeerIDs(ctx context.Context, transaction store.Store, accountID string, existing *types.NetworkRouter) []string { + var ids []string + if existing.Peer != "" { + ids = append(ids, existing.Peer) + } + if len(existing.PeerGroups) > 0 { + groupPeers, err := transaction.GetPeerIDsByGroups(ctx, accountID, existing.PeerGroups) + if err != nil { + log.WithContext(ctx).Errorf("failed to get old router peer-group members for affected peers: %v", err) + } else { + ids = append(ids, groupPeers...) + } + } + return ids } func (m *managerImpl) DeleteRouter(ctx context.Context, accountID, userID, networkID, routerID string) error { @@ -245,18 +257,9 @@ func (m *managerImpl) DeleteRouter(ctx context.Context, accountID, userID, netwo } var event func() - var affectedData *routerAffectedPeersData + var affectedPeerIDs []string 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) - } + affectedPeerIDs = m.accountManager.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{NetworkIDs: []string{networkID}}) event, err = m.DeleteRouterInTransaction(ctx, transaction, accountID, userID, networkID, routerID) if err != nil { @@ -276,7 +279,7 @@ func (m *managerImpl) DeleteRouter(ctx context.Context, accountID, userID, netwo event() - if affectedPeerIDs := m.resolveRouterAffectedPeers(ctx, accountID, affectedData); len(affectedPeerIDs) > 0 { + if 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 { @@ -313,153 +316,6 @@ 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 &routerAffectedPeersData{}, 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 - } - - 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 { - groupSet[gID] = struct{}{} - } - - if len(data.resourceGroupIDs) > 0 { - 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 := 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 { - 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 NewManagerMock() Manager { return &mockManager{} } diff --git a/management/server/peer.go b/management/server/peer.go index ad544aaa2..90fd0227a 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -27,6 +27,7 @@ import ( "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/affectedpeers" nbpeer "github.com/netbirdio/netbird/management/server/peer" "github.com/netbirdio/netbird/management/server/telemetry" "github.com/netbirdio/netbird/shared/management/status" @@ -1463,6 +1464,18 @@ func (am *DefaultAccountManager) BufferUpdateAffectedPeers(ctx context.Context, _ = am.networkMapController.BufferUpdateAffectedPeers(ctx, accountID, peerIDs, reason) } +// ResolveAffectedPeers resolves a description of what changed into the peer IDs +// whose network map may have changed. It is the single entry point shared by the +// server package and the networks managers (via the account.Manager interface). +func (am *DefaultAccountManager) ResolveAffectedPeers(ctx context.Context, s store.Store, accountID string, change affectedpeers.Change) []string { + peerIDs, err := affectedpeers.Resolve(ctx, s, accountID, change) + if err != nil { + log.WithContext(ctx).Errorf("failed to resolve affected peers: %v", err) + return nil + } + return peerIDs +} + // 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) @@ -1471,14 +1484,10 @@ func (am *DefaultAccountManager) resolveAffectedPeersForPeerChanges(ctx context. return nil } - log.WithContext(ctx).Tracef("resolveAffectedPeersForPeerChanges: changedPeers=%v -> groups=%v", changedPeerIDs, groupIDs) - - // 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)) - return result + return am.ResolveAffectedPeers(ctx, s, accountID, affectedpeers.Change{ + ChangedGroupIDs: groupIDs, + ChangedPeerIDs: changedPeerIDs, + }) } func (am *DefaultAccountManager) BufferUpdateAccountPeers(ctx context.Context, accountID string, reason types.UpdateReason) { diff --git a/management/server/policy.go b/management/server/policy.go index 5a5286088..d9b8598ae 100644 --- a/management/server/policy.go +++ b/management/server/policy.go @@ -13,6 +13,7 @@ import ( "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/affectedpeers" "github.com/netbirdio/netbird/management/server/posture" "github.com/netbirdio/netbird/shared/management/status" ) @@ -74,7 +75,7 @@ func (am *DefaultAccountManager) SavePolicy(ctx context.Context, accountID, user } } - affectedPeerIDs = am.resolvePolicyAffectedPeers(ctx, transaction, accountID, policy, existingPolicy) + affectedPeerIDs = am.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{Policies: []*types.Policy{policy, existingPolicy}}) return transaction.IncrementNetworkSerial(ctx, accountID) }) @@ -117,7 +118,7 @@ func (am *DefaultAccountManager) DeletePolicy(ctx context.Context, accountID, po return err } - affectedPeerIDs = am.resolvePolicyAffectedPeers(ctx, transaction, accountID, policy) + affectedPeerIDs = am.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{Policies: []*types.Policy{policy}}) if err = transaction.DeletePolicy(ctx, accountID, policyID); err != nil { return err @@ -154,43 +155,6 @@ func (am *DefaultAccountManager) ListPolicies(ctx context.Context, accountID, us return am.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) } -// collectPolicyAffectedGroupsAndPeers returns group IDs and direct peer IDs from the given policies. -func collectPolicyAffectedGroupsAndPeers(ctx context.Context, policies ...*types.Policy) (groupIDs []string, directPeerIDs []string) { - for _, policy := range policies { - if policy == nil { - continue - } - 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 -} - -// resolvePolicyAffectedPeers resolves the peers affected by the given policies into -// a deduplicated peer ID list. It combines the policies' literal rule groups and -// direct peers with the routing peers that serve any targeted network resource. -func (am *DefaultAccountManager) resolvePolicyAffectedPeers(ctx context.Context, transaction store.Store, accountID string, policies ...*types.Policy) []string { - groupIDs, directPeerIDs := collectPolicyAffectedGroupsAndPeers(ctx, policies...) - - groupSet := toSet(groupIDs) - peerSet := toSet(directPeerIDs) - collectPolicyRouterBridge(ctx, transaction, accountID, groupSet, peerSet, policies...) - - return am.resolvePeerIDs(ctx, transaction, accountID, setToSlice(groupSet), setToSlice(peerSet)) -} - // validatePolicy validates the policy and its rules. For updates it returns // the existing policy loaded from the store so callers can avoid a second read. func validatePolicy(ctx context.Context, transaction store.Store, accountID string, policy *types.Policy) (*types.Policy, error) { diff --git a/management/server/posture_checks.go b/management/server/posture_checks.go index cdfe169a2..f460145b5 100644 --- a/management/server/posture_checks.go +++ b/management/server/posture_checks.go @@ -8,6 +8,7 @@ import ( log "github.com/sirupsen/logrus" "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/affectedpeers" "github.com/netbirdio/netbird/management/server/permissions/modules" "github.com/netbirdio/netbird/management/server/permissions/operations" "github.com/netbirdio/netbird/management/server/posture" @@ -53,8 +54,7 @@ func (am *DefaultAccountManager) SavePostureChecks(ctx context.Context, accountI if isUpdate { action = activity.PostureCheckUpdated - groupIDs, directPeerIDs := collectPostureCheckAffectedGroupsAndPeers(ctx, transaction, accountID, postureChecks.ID) - affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) + affectedPeerIDs = am.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{PostureCheckIDs: []string{postureChecks.ID}}) } postureChecks.AccountID = accountID @@ -134,27 +134,6 @@ func (am *DefaultAccountManager) ListPostureChecks(ctx context.Context, accountI return am.Store.GetAccountPostureChecks(ctx, store.LockingStrengthNone, accountID) } -// 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 { - log.WithContext(ctx).Errorf("failed to get policies for posture check affected peers resolution: %v", err) - return nil, nil - } - - for _, policy := range policies { - if slices.Contains(policy.SourcePostureChecks, postureCheckID) { - 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 -} - // validatePostureChecks validates the posture checks. func validatePostureChecks(ctx context.Context, transaction store.Store, accountID string, postureChecks *posture.Checks) error { if err := postureChecks.Validate(); err != nil { diff --git a/management/server/route.go b/management/server/route.go index 075d56ba6..44f2528b7 100644 --- a/management/server/route.go +++ b/management/server/route.go @@ -11,6 +11,7 @@ import ( log "github.com/sirupsen/logrus" "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/affectedpeers" "github.com/netbirdio/netbird/management/server/permissions/modules" "github.com/netbirdio/netbird/management/server/permissions/operations" "github.com/netbirdio/netbird/management/server/store" @@ -178,8 +179,7 @@ func (am *DefaultAccountManager) CreateRoute(ctx context.Context, accountID stri return err } - groupIDs, directPeerIDs := collectRouteAffectedGroupsAndPeers(ctx, newRoute) - affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) + affectedPeerIDs = am.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{Routes: []*route.Route{newRoute}}) return transaction.IncrementNetworkSerial(ctx, accountID) }) @@ -228,8 +228,7 @@ func (am *DefaultAccountManager) SaveRoute(ctx context.Context, accountID, userI return err } - groupIDs, directPeerIDs := collectRouteAffectedGroupsAndPeers(ctx, routeToSave, oldRoute) - affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) + affectedPeerIDs = am.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{Routes: []*route.Route{routeToSave, oldRoute}}) return transaction.IncrementNetworkSerial(ctx, accountID) }) @@ -268,8 +267,7 @@ func (am *DefaultAccountManager) DeleteRoute(ctx context.Context, accountID stri return err } - groupIDs, directPeerIDs := collectRouteAffectedGroupsAndPeers(ctx, rt) - affectedPeerIDs = am.resolvePeerIDs(ctx, transaction, accountID, groupIDs, directPeerIDs) + affectedPeerIDs = am.ResolveAffectedPeers(ctx, transaction, accountID, affectedpeers.Change{Routes: []*route.Route{rt}}) if err = transaction.DeleteRoute(ctx, accountID, string(routeID)); err != nil { return err @@ -376,25 +374,6 @@ func getPlaceholderIP() netip.Prefix { return netip.PrefixFrom(netip.AddrFrom4([4]byte{192, 0, 2, 0}), 32) } -// collectRouteAffectedGroupsAndPeers returns group IDs and direct peer IDs from the given routes. -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...) - if r.Peer != "" { - directPeerIDs = append(directPeerIDs, r.Peer) - } - } - log.WithContext(ctx).Tracef("collectRouteAffectedGroupsAndPeers: result groupIDs=%v, directPeerIDs=%v", groupIDs, directPeerIDs) - return -} - // GetRoutesByPrefixOrDomains return list of routes by account and route prefix func getRoutesByPrefixOrDomains(ctx context.Context, transaction store.Store, accountID string, prefix netip.Prefix, domains domain.List) ([]*route.Route, error) { accountRoutes, err := transaction.GetAccountRoutes(ctx, store.LockingStrengthNone, accountID)