diff --git a/management/server/networks/manager.go b/management/server/networks/manager.go index c96b60bb2..8da0c594d 100644 --- a/management/server/networks/manager.go +++ b/management/server/networks/manager.go @@ -71,9 +71,20 @@ func (m *managerImpl) CreateNetwork(ctx context.Context, userID string, network network.ID = xid.New().String() - err = m.store.SaveNetwork(ctx, network) + err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { + seq, err := transaction.AllocateAccountSeqID(ctx, network.AccountID, serverTypes.AccountSeqEntityNetwork) + if err != nil { + return fmt.Errorf("failed to allocate network seq id: %w", err) + } + network.AccountSeqID = seq + + if err := transaction.SaveNetwork(ctx, network); err != nil { + return fmt.Errorf("failed to save network: %w", err) + } + return nil + }) if err != nil { - return nil, fmt.Errorf("failed to save network: %w", err) + return nil, err } m.accountManager.StoreEvent(ctx, userID, network.ID, network.AccountID, activity.NetworkCreated, network.EventMeta()) @@ -102,14 +113,25 @@ func (m *managerImpl) UpdateNetwork(ctx context.Context, userID string, network return nil, status.NewPermissionDeniedError() } - _, err = m.store.GetNetworkByID(ctx, store.LockingStrengthUpdate, network.AccountID, network.ID) + err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { + existing, err := transaction.GetNetworkByID(ctx, store.LockingStrengthUpdate, network.AccountID, network.ID) + if err != nil { + return fmt.Errorf("failed to get network: %w", err) + } + network.AccountSeqID = existing.AccountSeqID + + if err := transaction.SaveNetwork(ctx, network); err != nil { + return fmt.Errorf("failed to save network: %w", err) + } + return nil + }) if err != nil { - return nil, fmt.Errorf("failed to get network: %w", err) + return nil, err } m.accountManager.StoreEvent(ctx, userID, network.ID, network.AccountID, activity.NetworkUpdated, network.EventMeta()) - return network, m.store.SaveNetwork(ctx, network) + return network, nil } func (m *managerImpl) DeleteNetwork(ctx context.Context, accountID, userID, networkID string) error { diff --git a/management/server/networks/manager_test.go b/management/server/networks/manager_test.go index 6fb19d157..a86d7e4e6 100644 --- a/management/server/networks/manager_test.go +++ b/management/server/networks/manager_test.go @@ -252,3 +252,73 @@ func Test_UpdateNetworkFailsWithPermissionDenied(t *testing.T) { require.Error(t, err) require.Nil(t, updatedNetwork) } + +// Test_CreateNetworkAllocatesSeqID verifies that CreateNetwork sets a +// non-zero AccountSeqID on the persisted network (allocated through the +// account_seq_counters table). +func Test_CreateNetworkAllocatesSeqID(t *testing.T) { + ctx := context.Background() + const accountID = "testAccountId" + const userID = "testAdminId" + + s, cleanUp, err := store.NewTestStoreFromSQL(ctx, "../testdata/networks.sql", t.TempDir()) + require.NoError(t, err) + t.Cleanup(cleanUp) + + am := mock_server.MockAccountManager{} + permissionsManager := permissions.NewManager(s) + groupsManager := groups.NewManagerMock() + routerManager := routers.NewManagerMock() + resourcesManager := resources.NewManager(s, permissionsManager, groupsManager, &am, nil) + manager := NewManager(s, permissionsManager, resourcesManager, routerManager, &am) + + created, err := manager.CreateNetwork(ctx, userID, &types.Network{ + AccountID: accountID, + Name: "seq-allocation-test", + }) + require.NoError(t, err) + require.NotZero(t, created.AccountSeqID, "CreateNetwork must allocate a non-zero AccountSeqID") +} + +// Test_UpdateNetworkPreservesSeqID verifies UpdateNetwork does not reset +// AccountSeqID even when the caller passes a zero value (the shape REST +// handlers produce because the field is `json:"-"`). +func Test_UpdateNetworkPreservesSeqID(t *testing.T) { + ctx := context.Background() + const accountID = "testAccountId" + const userID = "testAdminId" + + s, cleanUp, err := store.NewTestStoreFromSQL(ctx, "../testdata/networks.sql", t.TempDir()) + require.NoError(t, err) + t.Cleanup(cleanUp) + + am := mock_server.MockAccountManager{} + permissionsManager := permissions.NewManager(s) + groupsManager := groups.NewManagerMock() + routerManager := routers.NewManagerMock() + resourcesManager := resources.NewManager(s, permissionsManager, groupsManager, &am, nil) + manager := NewManager(s, permissionsManager, resourcesManager, routerManager, &am) + + created, err := manager.CreateNetwork(ctx, userID, &types.Network{ + AccountID: accountID, + Name: "seq-preserve-original", + }) + require.NoError(t, err) + originalSeq := created.AccountSeqID + require.NotZero(t, originalSeq) + + update := &types.Network{ + AccountID: accountID, + ID: created.ID, + Name: "seq-preserve-renamed", + } + require.Zero(t, update.AccountSeqID, "incoming struct must mirror an HTTP handler shape") + + _, err = manager.UpdateNetwork(ctx, userID, update) + require.NoError(t, err) + + got, err := manager.GetNetwork(ctx, accountID, userID, created.ID) + require.NoError(t, err) + require.Equal(t, originalSeq, got.AccountSeqID, "AccountSeqID must survive UpdateNetwork") + require.Equal(t, "seq-preserve-renamed", got.Name) +} diff --git a/management/server/networks/types/network.go b/management/server/networks/types/network.go index 69d596f8b..0c6920879 100644 --- a/management/server/networks/types/network.go +++ b/management/server/networks/types/network.go @@ -7,12 +7,24 @@ import ( ) type Network struct { - ID string `gorm:"primaryKey"` - AccountID string `gorm:"index"` + ID string `gorm:"primaryKey"` + AccountID string `gorm:"index"` + + // AccountSeqID is a per-account monotonically increasing identifier used as the + // compact wire id when sending NetworkMap components to capable peers. + AccountSeqID uint32 `json:"-" gorm:"index:idx_networks_account_seq_id;not null;default:0"` + Name string Description string } +// HasSeqID reports whether the network has been persisted long enough to have +// a per-account sequence id allocated. Wire encoders that key off AccountSeqID +// must skip networks that return false here. +func (n *Network) HasSeqID() bool { + return n != nil && n.AccountSeqID != 0 +} + func NewNetwork(accountId, name, description string) *Network { return &Network{ ID: xid.New().String(), @@ -41,13 +53,14 @@ func (n *Network) FromAPIRequest(req *api.NetworkRequest) { } } -// Copy returns a copy of a posture checks. +// Copy returns a copy of a network. func (n *Network) Copy() *Network { return &Network{ - ID: n.ID, - AccountID: n.AccountID, - Name: n.Name, - Description: n.Description, + ID: n.ID, + AccountID: n.AccountID, + AccountSeqID: n.AccountSeqID, + Name: n.Name, + Description: n.Description, } } diff --git a/management/server/posture/checks.go b/management/server/posture/checks.go index f0bbbc32e..a79cf6898 100644 --- a/management/server/posture/checks.go +++ b/management/server/posture/checks.go @@ -47,10 +47,21 @@ type Checks struct { // AccountID is a reference to the Account that this object belongs AccountID string `json:"-" gorm:"index"` + // AccountSeqID is a per-account monotonically increasing identifier used as the + // compact wire id when sending NetworkMap components to capable peers. + AccountSeqID uint32 `json:"-" gorm:"index:idx_posture_checks_account_seq_id;not null;default:0"` + // Checks is a set of objects that perform the actual checks Checks ChecksDefinition `gorm:"serializer:json"` } +// HasSeqID reports whether the posture check has been persisted long enough +// to have a per-account sequence id allocated. Wire encoders that key off +// AccountSeqID must skip checks that return false here. +func (pc *Checks) HasSeqID() bool { + return pc != nil && pc.AccountSeqID != 0 +} + // ChecksDefinition contains definition of actual check type ChecksDefinition struct { NBVersionCheck *NBVersionCheck `json:",omitempty"` @@ -121,11 +132,12 @@ func (*Checks) TableName() string { // Copy returns a copy of a posture checks. func (pc *Checks) Copy() *Checks { checks := &Checks{ - ID: pc.ID, - Name: pc.Name, - Description: pc.Description, - AccountID: pc.AccountID, - Checks: pc.Checks.Copy(), + ID: pc.ID, + Name: pc.Name, + Description: pc.Description, + AccountID: pc.AccountID, + AccountSeqID: pc.AccountSeqID, + Checks: pc.Checks.Copy(), } return checks } diff --git a/management/server/posture_checks.go b/management/server/posture_checks.go index 1e3ce4b8a..5f548f2de 100644 --- a/management/server/posture_checks.go +++ b/management/server/posture_checks.go @@ -51,12 +51,24 @@ func (am *DefaultAccountManager) SavePostureChecks(ctx context.Context, accountI } if isUpdate { + existing, err := transaction.GetPostureChecksByID(ctx, store.LockingStrengthNone, accountID, postureChecks.ID) + if err != nil { + return err + } + postureChecks.AccountSeqID = existing.AccountSeqID + updateAccountPeers, err = arePostureCheckChangesAffectPeers(ctx, transaction, accountID, postureChecks.ID) if err != nil { return err } action = activity.PostureCheckUpdated + } else { + seq, err := transaction.AllocateAccountSeqID(ctx, accountID, types.AccountSeqEntityPostureCheck) + if err != nil { + return err + } + postureChecks.AccountSeqID = seq } postureChecks.AccountID = accountID diff --git a/management/server/posture_checks_test.go b/management/server/posture_checks_test.go index 394f0d896..ca9bef389 100644 --- a/management/server/posture_checks_test.go +++ b/management/server/posture_checks_test.go @@ -563,3 +563,61 @@ func TestArePostureCheckChangesAffectPeers(t *testing.T) { assert.False(t, result) }) } + +// TestSavePostureChecks_AllocatesSeqIDOnCreate verifies that the create path +// (no incoming ID) allocates a non-zero AccountSeqID via the +// account_seq_counters table. +func TestSavePostureChecks_AllocatesSeqIDOnCreate(t *testing.T) { + am, _, err := createManager(t) + require.NoError(t, err) + + account, err := initTestPostureChecksAccount(am) + require.NoError(t, err) + + created, err := am.SavePostureChecks(context.Background(), account.Id, adminUserID, &posture.Checks{ + Name: "seq-allocation-test", + Checks: posture.ChecksDefinition{ + NBVersionCheck: &posture.NBVersionCheck{MinVersion: "0.26.0"}, + }, + }, true) + require.NoError(t, err) + require.NotZero(t, created.AccountSeqID, "SavePostureChecks on create must allocate a non-zero AccountSeqID") +} + +// TestSavePostureChecks_PreservesSeqIDOnUpdate verifies the update path does +// not reset AccountSeqID even when the caller passes a zero value (REST +// handler shape, because the field is `json:"-"`). +func TestSavePostureChecks_PreservesSeqIDOnUpdate(t *testing.T) { + am, _, err := createManager(t) + require.NoError(t, err) + + account, err := initTestPostureChecksAccount(am) + require.NoError(t, err) + + created, err := am.SavePostureChecks(context.Background(), account.Id, adminUserID, &posture.Checks{ + Name: "seq-preserve-original", + Checks: posture.ChecksDefinition{ + NBVersionCheck: &posture.NBVersionCheck{MinVersion: "0.26.0"}, + }, + }, true) + require.NoError(t, err) + originalSeq := created.AccountSeqID + require.NotZero(t, originalSeq) + + update := &posture.Checks{ + ID: created.ID, + Name: "seq-preserve-renamed", + Checks: posture.ChecksDefinition{ + NBVersionCheck: &posture.NBVersionCheck{MinVersion: "0.27.0"}, + }, + } + require.Zero(t, update.AccountSeqID, "incoming struct must mirror an HTTP handler shape") + + _, err = am.SavePostureChecks(context.Background(), account.Id, adminUserID, update, false) + require.NoError(t, err) + + got, err := am.GetPostureChecks(context.Background(), account.Id, created.ID, adminUserID) + require.NoError(t, err) + require.Equal(t, originalSeq, got.AccountSeqID, "AccountSeqID must survive SavePostureChecks update") + require.Equal(t, "seq-preserve-renamed", got.Name) +} diff --git a/management/server/store/account_seq_test.go b/management/server/store/account_seq_test.go index 4a8134559..7ee52e4cc 100644 --- a/management/server/store/account_seq_test.go +++ b/management/server/store/account_seq_test.go @@ -11,6 +11,8 @@ import ( nbdns "github.com/netbirdio/netbird/dns" resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" + networkTypes "github.com/netbirdio/netbird/management/server/networks/types" + "github.com/netbirdio/netbird/management/server/posture" "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/route" ) @@ -208,6 +210,7 @@ func TestSaveAccount_PreservesExistingSeqIDs(t *testing.T) { nsgSeqs := make(map[string]uint32) resourceSeqs := make(map[string]uint32) routerSeqs := make(map[string]uint32) + networkSeqs := make(map[string]uint32) for _, g := range account.Groups { require.NotZero(t, g.AccountSeqID, "fixture group must have seq>0 after backfill") @@ -233,6 +236,10 @@ func TestSaveAccount_PreservesExistingSeqIDs(t *testing.T) { require.NotZero(t, nr.AccountSeqID, "fixture network_router must have seq>0") routerSeqs[nr.ID] = nr.AccountSeqID } + for _, n := range account.Networks { + require.NotZero(t, n.AccountSeqID, "fixture network must have seq>0 after backfill") + networkSeqs[n.ID] = n.AccountSeqID + } require.NoError(t, store.SaveAccount(ctx, account)) @@ -256,6 +263,9 @@ func TestSaveAccount_PreservesExistingSeqIDs(t *testing.T) { for _, nr := range after.NetworkRouters { require.Equal(t, routerSeqs[nr.ID], nr.AccountSeqID, "network_router %s seq must be preserved", nr.ID) } + for _, n := range after.Networks { + require.Equal(t, networkSeqs[n.ID], n.AccountSeqID, "network %s seq must be preserved", n.ID) + } } func TestSaveAccount_AllocatesSeqIDsForAllEntityTypes(t *testing.T) { @@ -298,6 +308,15 @@ func TestSaveAccount_AllocatesSeqIDsForAllEntityTypes(t *testing.T) { NetworkRouters: []*routerTypes.NetworkRouter{ {ID: "nrt1", AccountID: accountID, NetworkID: "net1", Peer: "peer1", Enabled: true}, }, + Networks: []*networkTypes.Network{ + {ID: "n1", AccountID: accountID, Name: "n1"}, + }, + PostureChecks: []*posture.Checks{ + {ID: "pc1", AccountID: accountID, Name: "pc1", + Checks: posture.ChecksDefinition{ + NBVersionCheck: &posture.NBVersionCheck{MinVersion: "0.26.0"}, + }}, + }, } require.NoError(t, store.SaveAccount(ctx, account)) @@ -311,6 +330,8 @@ func TestSaveAccount_AllocatesSeqIDsForAllEntityTypes(t *testing.T) { require.Len(t, after.NameServerGroups, 1) require.Len(t, after.NetworkResources, 1) require.Len(t, after.NetworkRouters, 1) + require.Len(t, after.Networks, 1) + require.Len(t, after.PostureChecks, 1) for _, g := range after.Groups { require.NotZero(t, g.AccountSeqID, "group seq must be allocated") @@ -330,6 +351,12 @@ func TestSaveAccount_AllocatesSeqIDsForAllEntityTypes(t *testing.T) { for _, nr := range after.NetworkRouters { require.NotZero(t, nr.AccountSeqID, "network_router seq must be allocated") } + for _, n := range after.Networks { + require.NotZero(t, n.AccountSeqID, "network seq must be allocated") + } + for _, pc := range after.PostureChecks { + require.NotZero(t, pc.AccountSeqID, "posture_check seq must be allocated") + } require.NoError(t, store.SaveAccount(ctx, after)) final, err := store.GetAccount(ctx, accountID) @@ -340,6 +367,20 @@ func TestSaveAccount_AllocatesSeqIDsForAllEntityTypes(t *testing.T) { for _, n := range final.NameServerGroups { require.Equal(t, after.NameServerGroups[n.ID].AccountSeqID, n.AccountSeqID, "name_server_group seq preserved on re-save") } + afterByID := map[string]uint32{} + for _, n := range after.Networks { + afterByID[n.ID] = n.AccountSeqID + } + for _, n := range final.Networks { + require.Equal(t, afterByID[n.ID], n.AccountSeqID, "network seq preserved on re-save") + } + afterPCByID := map[string]uint32{} + for _, pc := range after.PostureChecks { + afterPCByID[pc.ID] = pc.AccountSeqID + } + for _, pc := range final.PostureChecks { + require.Equal(t, afterPCByID[pc.ID], pc.AccountSeqID, "posture_check seq preserved on re-save") + } } func TestAllocateAccountSeqID_ConcurrentSameAccountEntity(t *testing.T) { diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go index be0e7f216..855e54af8 100644 --- a/management/server/store/sql_store.go +++ b/management/server/store/sql_store.go @@ -3724,6 +3724,26 @@ func (s *SqlStore) assignAccountSeqIDs(ctx context.Context, tx *gorm.DB, account } nr.AccountSeqID = seq } + for _, n := range account.Networks { + if n == nil || n.AccountSeqID != 0 { + continue + } + seq, err := allocateAccountSeqID(ctx, tx, s.storeEngine, account.Id, types.AccountSeqEntityNetwork) + if err != nil { + return err + } + n.AccountSeqID = seq + } + for _, pc := range account.PostureChecks { + if pc == nil || pc.AccountSeqID != 0 { + continue + } + seq, err := allocateAccountSeqID(ctx, tx, s.storeEngine, account.Id, types.AccountSeqEntityPostureCheck) + if err != nil { + return err + } + pc.AccountSeqID = seq + } return nil } diff --git a/management/server/store/store.go b/management/server/store/store.go index 28aa2e264..9b4ecd8b0 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -545,6 +545,12 @@ func getMigrationsPostAuto(ctx context.Context) []migrationFunc { func(db *gorm.DB) error { return migration.BackfillAccountSeqIDs[dns.NameServerGroup](ctx, db, types.AccountSeqEntityNameserverGroup, "id") }, + func(db *gorm.DB) error { + return migration.BackfillAccountSeqIDs[networkTypes.Network](ctx, db, types.AccountSeqEntityNetwork, "id") + }, + func(db *gorm.DB) error { + return migration.BackfillAccountSeqIDs[posture.Checks](ctx, db, types.AccountSeqEntityPostureCheck, "id") + }, } } diff --git a/management/server/types/account_seq_counter.go b/management/server/types/account_seq_counter.go index e21f73317..80ce74099 100644 --- a/management/server/types/account_seq_counter.go +++ b/management/server/types/account_seq_counter.go @@ -10,6 +10,8 @@ const ( AccountSeqEntityNetworkResource AccountSeqEntity = "network_resource" AccountSeqEntityNetworkRouter AccountSeqEntity = "network_router" AccountSeqEntityNameserverGroup AccountSeqEntity = "nameserver_group" + AccountSeqEntityNetwork AccountSeqEntity = "network" + AccountSeqEntityPostureCheck AccountSeqEntity = "posture_check" ) // AccountSeqCounter tracks the next per-account integer id for a given component