diff --git a/management/server/account.go b/management/server/account.go index 619036d0c..617231b46 100644 --- a/management/server/account.go +++ b/management/server/account.go @@ -33,6 +33,7 @@ import ( nbconfig "github.com/netbirdio/netbird/management/internals/server/config" "github.com/netbirdio/netbird/management/server/account" "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/affectedpeers" nbcache "github.com/netbirdio/netbird/management/server/cache" nbcontext "github.com/netbirdio/netbird/management/server/context" "github.com/netbirdio/netbird/management/server/geolocation" @@ -1626,6 +1627,8 @@ func (am *DefaultAccountManager) SyncUserJWTGroups(ctx context.Context, userAuth var removeOldGroups []string var hasChanges bool var user *types.User + var change affectedpeers.Change + var snap *affectedpeers.Snapshot err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { user, err = transaction.GetUserByUserID(ctx, store.LockingStrengthNone, userAuth.UserId) if err != nil { @@ -1664,6 +1667,8 @@ func (am *DefaultAccountManager) SyncUserJWTGroups(ctx context.Context, userAuth return fmt.Errorf("error saving user: %w", err) } + allGroupChanges := slices.Concat(addNewGroups, removeOldGroups) + // Propagate changes to peers if group propagation is enabled if settings.GroupsPropagationEnabled { peers, err := transaction.GetUserPeers(ctx, store.LockingStrengthNone, userAuth.AccountId, userAuth.UserId) @@ -1672,6 +1677,7 @@ func (am *DefaultAccountManager) SyncUserJWTGroups(ctx context.Context, userAuth } for _, peer := range peers { + change.OutputPeerIDs = append(change.OutputPeerIDs, peer.ID) for _, g := range addNewGroups { if err := transaction.AddPeerToGroup(ctx, userAuth.AccountId, peer.ID, g); err != nil { return fmt.Errorf("error adding peer %s to group %s: %w", peer.ID, g, err) @@ -1684,7 +1690,19 @@ func (am *DefaultAccountManager) SyncUserJWTGroups(ctx context.Context, userAuth } } - allGroupChanges := slices.Concat(addNewGroups, removeOldGroups) + change.LinkGroups = allGroupChanges + + // An IPv6 reconcile can change the peers' addresses, which are visible + // through ALL their group memberships, so seed the full walk instead of + // folding the peers alone. + for _, g := range allGroupChanges { + if slices.Contains(settings.IPv6EnabledGroups, g) { + change.ChangedPeerIDs = change.OutputPeerIDs + change.OutputPeerIDs = nil + break + } + } + if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, userAuth.AccountId, allGroupChanges); err != nil { return fmt.Errorf("reconcile IPv6 for group changes: %w", err) } @@ -1694,6 +1712,20 @@ func (am *DefaultAccountManager) SyncUserJWTGroups(ctx context.Context, userAuth } } + // The user->group mapping shipped to SSH destination peers (GroupIDToUserIDs) + // changed for these groups even without peer propagation, so the destinations + // of SSH rules authorizing them must refresh. + sshDistributionGroups, sshPeerIDs, err := sshAuthorizedGroupConsumers(ctx, transaction, userAuth.AccountId, allGroupChanges) + if err != nil { + return err + } + change.DistributionGroupIDs = sshDistributionGroups + change.OutputPeerIDs = append(change.OutputPeerIDs, sshPeerIDs...) + + if snap, err = affectedpeers.Load(ctx, transaction, userAuth.AccountId, change); err != nil { + return err + } + return nil }) if err != nil { @@ -1730,24 +1762,68 @@ func (am *DefaultAccountManager) SyncUserJWTGroups(ctx context.Context, userAuth } } - removedGroupAffectsPeers, err := areGroupChangesAffectPeers(ctx, am.Store, userAuth.AccountId, removeOldGroups) - if err != nil { - return err - } - - newGroupsAffectsPeers, err := areGroupChangesAffectPeers(ctx, am.Store, userAuth.AccountId, addNewGroups) - if err != nil { - return err - } - - if removedGroupAffectsPeers || newGroupsAffectsPeers { - log.WithContext(ctx).Tracef("user %s: JWT group membership changed, updating account peers", userAuth.UserId) - am.BufferUpdateAccountPeers(ctx, userAuth.AccountId, types.UpdateReason{Resource: types.UpdateResourceUser, Operation: types.UpdateOperationUpdate}) - } + log.WithContext(ctx).Tracef("user %s: JWT group membership changed, updating affected peers", userAuth.UserId) + bgCtx := context.WithoutCancel(ctx) + go func() { + affectedPeerIDs := snap.Expand(bgCtx, userAuth.AccountId, change) + if len(affectedPeerIDs) == 0 { + return + } + if err := am.networkMapController.BufferUpdateAffectedPeers(bgCtx, userAuth.AccountId, affectedPeerIDs, types.UpdateReason{Resource: types.UpdateResourceUser, Operation: types.UpdateOperationUpdate}); err != nil { + log.WithContext(bgCtx).Errorf("failed to update affected peers after JWT group sync for account %s: %v", userAuth.AccountId, err) + } + }() return nil } +// sshAuthorizedGroupConsumers returns the destination groups and destination peers of +// enabled SSH rules whose AuthorizedGroups reference any of the changed groups. Their +// network maps carry the user->group mapping for those groups. +func sshAuthorizedGroupConsumers(ctx context.Context, transaction store.Store, accountID string, changedGroupIDs []string) ([]string, []string, error) { + if len(changedGroupIDs) == 0 { + return nil, nil, nil + } + + policies, err := transaction.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + if err != nil { + return nil, nil, fmt.Errorf("error getting account policies: %w", err) + } + + changed := make(map[string]struct{}, len(changedGroupIDs)) + for _, id := range changedGroupIDs { + changed[id] = struct{}{} + } + + var groupIDs []string + var peerIDs []string + for _, policy := range policies { + if policy == nil || !policy.Enabled { + continue + } + for _, rule := range policy.Rules { + if !rule.Enabled || rule.Protocol != types.PolicyRuleProtocolNetbirdSSH { + continue + } + authorized := false + for groupID := range rule.AuthorizedGroups { + if _, ok := changed[groupID]; ok { + authorized = true + break + } + } + if !authorized { + continue + } + groupIDs = append(groupIDs, rule.Destinations...) + if rule.DestinationResource.Type == types.ResourceTypePeer && rule.DestinationResource.ID != "" { + peerIDs = append(peerIDs, rule.DestinationResource.ID) + } + } + } + return groupIDs, peerIDs, nil +} + // getAccountIDWithAuthorizationClaims retrieves an account ID using JWT Claims. // if domain is not private or domain is invalid, it will return the account ID by user ID. // if domain is of the PrivateCategory category, it will evaluate diff --git a/management/server/affected_peers_jwt_test.go b/management/server/affected_peers_jwt_test.go new file mode 100644 index 000000000..32e54797d --- /dev/null +++ b/management/server/affected_peers_jwt_test.go @@ -0,0 +1,110 @@ +package server + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" + + nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/auth" +) + +// TestAffectedPeers_SyncUserJWTGroups_OnlyAffectedPeersUpdated verifies that a JWT +// auto-group change updates only the user's peers and the peers linked to the changed +// group through policies, instead of fanning out to the whole account. +func TestAffectedPeers_SyncUserJWTGroups_OnlyAffectedPeersUpdated(t *testing.T) { + manager, updateManager, account, _, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + key, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err) + userPeer, _, _, _, err := manager.AddPeer(ctx, accountID, "", userID, &nbpeer.Peer{ + Key: key.PublicKey().String(), + Meta: nbpeer.PeerSystemMeta{Hostname: "user-peer"}, + }, false) + require.NoError(t, err) + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + require.NoError(t, manager.Store.DeletePolicy(ctx, accountID, p.ID)) + } + + account, err = manager.Store.GetAccount(ctx, accountID) + require.NoError(t, err) + account.Settings.JWTGroupsEnabled = true + account.Settings.JWTGroupsClaimName = "groups" + account.Settings.GroupsPropagationEnabled = true + require.NoError(t, manager.Store.SaveAccount(ctx, account)) + + require.NoError(t, manager.CreateGroup(ctx, accountID, userID, &types.Group{ID: "jwt-grp", Name: "jwt-linked", Issued: types.GroupIssuedJWT, Peers: []string{}})) + require.NoError(t, manager.CreateGroup(ctx, accountID, userID, &types.Group{ID: "jwt-dest", Name: "jwt-dest", Peers: []string{peer2.ID}})) + + _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{"jwt-grp"}, + Destinations: []string{"jwt-dest"}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + }, + }, true) + require.NoError(t, err) + + updUser := updateManager.CreateChannel(ctx, userPeer.ID) + upd2 := updateManager.CreateChannel(ctx, peer2.ID) + upd3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, userPeer.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + userAuth := auth.UserAuth{ + AccountId: accountID, + UserId: userID, + Groups: []string{"jwt-linked"}, + } + + t.Run("adding JWT group updates only linked peers", func(t *testing.T) { + drainPeerUpdates(updUser) + drainPeerUpdates(upd2) + drainPeerUpdates(upd3) + + require.NoError(t, manager.SyncUserJWTGroups(ctx, userAuth)) + + peerShouldReceiveUpdate(t, updUser) + peerShouldReceiveUpdate(t, upd2) + peerShouldNotReceiveUpdate(t, upd3) + + user, err := manager.Store.GetUserByUserID(ctx, store.LockingStrengthNone, userID) + require.NoError(t, err) + assert.Contains(t, user.AutoGroups, "jwt-grp") + }) + + t.Run("removing JWT group updates only linked peers", func(t *testing.T) { + drainPeerUpdates(updUser) + drainPeerUpdates(upd2) + drainPeerUpdates(upd3) + + userAuth.Groups = nil + require.NoError(t, manager.SyncUserJWTGroups(ctx, userAuth)) + + peerShouldReceiveUpdate(t, updUser) + peerShouldReceiveUpdate(t, upd2) + peerShouldNotReceiveUpdate(t, upd3) + + user, err := manager.Store.GetUserByUserID(ctx, store.LockingStrengthNone, userID) + require.NoError(t, err) + assert.NotContains(t, user.AutoGroups, "jwt-grp") + }) +}