diff --git a/management/internals/modules/zones/manager/manager.go b/management/internals/modules/zones/manager/manager.go index d5348d3d0..6f6ba6c40 100644 --- a/management/internals/modules/zones/manager/manager.go +++ b/management/internals/modules/zones/manager/manager.go @@ -3,15 +3,16 @@ package manager import ( "context" "fmt" + "slices" "github.com/netbirdio/netbird/management/internals/modules/zones" "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/permissions" "github.com/netbirdio/netbird/management/server/permissions/modules" "github.com/netbirdio/netbird/management/server/permissions/operations" "github.com/netbirdio/netbird/management/server/store" - "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/shared/management/status" ) @@ -69,6 +70,9 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string, } zone = zones.NewZone(accountID, zone.Name, zone.Domain, zone.Enabled, zone.EnableSearchDomain, zone.DistributionGroups) + var snap *affectedpeers.Snapshot + change := affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups} + err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { existingZone, err := transaction.GetZoneByDomain(ctx, accountID, zone.Domain) if err != nil { @@ -88,7 +92,15 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string, } if err = transaction.CreateZone(ctx, zone); err != nil { - return fmt.Errorf("failed to create zone: %w", err) + return fmt.Errorf("create zone: %w", err) + } + + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + + if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil { + return fmt.Errorf("increment network serial: %w", err) } return nil @@ -99,6 +111,8 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string, m.accountManager.StoreEvent(ctx, userID, zone.ID, accountID, activity.DNSZoneCreated, zone.EventMeta()) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) + return zone, nil } @@ -111,21 +125,26 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string, return nil, status.NewPermissionDeniedError() } - zone, err := m.store.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, updatedZone.ID) - if err != nil { - return nil, fmt.Errorf("failed to get zone: %w", err) - } - - if zone.Domain != updatedZone.Domain { - return nil, status.Errorf(status.InvalidArgument, "zone domain cannot be updated") - } - - zone.Name = updatedZone.Name - zone.Enabled = updatedZone.Enabled - zone.EnableSearchDomain = updatedZone.EnableSearchDomain - zone.DistributionGroups = updatedZone.DistributionGroups + var zone *zones.Zone + var snap *affectedpeers.Snapshot + var change affectedpeers.Change err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { + zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, updatedZone.ID) + if err != nil { + return fmt.Errorf("get zone: %w", err) + } + + if zone.Domain != updatedZone.Domain { + return status.Errorf(status.InvalidArgument, "zone domain cannot be updated") + } + + oldGroups := zone.DistributionGroups + zone.Name = updatedZone.Name + zone.Enabled = updatedZone.Enabled + zone.EnableSearchDomain = updatedZone.EnableSearchDomain + zone.DistributionGroups = updatedZone.DistributionGroups + for _, groupID := range zone.DistributionGroups { _, err = transaction.GetGroupByID(ctx, store.LockingStrengthNone, accountID, groupID) if err != nil { @@ -134,7 +153,16 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string, } if err = transaction.UpdateZone(ctx, zone); err != nil { - return fmt.Errorf("failed to update zone: %w", err) + return fmt.Errorf("update zone: %w", err) + } + + change = affectedpeers.Change{DistributionGroupIDs: slices.Concat(zone.DistributionGroups, oldGroups)} + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + + if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil { + return fmt.Errorf("increment network serial: %w", err) } return nil @@ -145,7 +173,7 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string, m.accountManager.StoreEvent(ctx, userID, zone.ID, accountID, activity.DNSZoneUpdated, zone.EventMeta()) - go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZone, Operation: types.UpdateOperationUpdate}) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) return zone, nil } @@ -159,13 +187,23 @@ func (m *managerImpl) DeleteZone(ctx context.Context, accountID, userID, zoneID return status.NewPermissionDeniedError() } - zone, err := m.store.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID) - if err != nil { - return fmt.Errorf("failed to get zone: %w", err) - } - + var zone *zones.Zone + var snap *affectedpeers.Snapshot + var change affectedpeers.Change var eventsToStore []func() + err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { + zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID) + if err != nil { + return fmt.Errorf("get zone: %w", err) + } + + // Load before delete: the post-delete state no longer references the groups. + change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups} + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + records, err := transaction.GetZoneDNSRecords(ctx, store.LockingStrengthNone, accountID, zoneID) if err != nil { return fmt.Errorf("failed to get records: %w", err) @@ -207,7 +245,7 @@ func (m *managerImpl) DeleteZone(ctx context.Context, accountID, userID, zoneID event() } - go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZone, Operation: types.UpdateOperationDelete}) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) return nil } diff --git a/management/internals/modules/zones/records/manager/manager.go b/management/internals/modules/zones/records/manager/manager.go index b041aca30..16839c1b4 100644 --- a/management/internals/modules/zones/records/manager/manager.go +++ b/management/internals/modules/zones/records/manager/manager.go @@ -9,11 +9,11 @@ import ( "github.com/netbirdio/netbird/management/internals/modules/zones/records" "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/permissions" "github.com/netbirdio/netbird/management/server/permissions/modules" "github.com/netbirdio/netbird/management/server/permissions/operations" "github.com/netbirdio/netbird/management/server/store" - "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/shared/management/status" ) @@ -65,6 +65,8 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI } var zone *zones.Zone + var snap *affectedpeers.Snapshot + var change affectedpeers.Change record = records.NewRecord(accountID, zoneID, record.Name, record.Type, record.Content, record.TTL) err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { @@ -82,6 +84,11 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI return fmt.Errorf("failed to create dns record: %w", err) } + change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups} + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + err = transaction.IncrementNetworkSerial(ctx, accountID) if err != nil { return fmt.Errorf("failed to increment network serial: %w", err) @@ -96,7 +103,7 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI meta := record.EventMeta(zone.ID, zone.Name) m.accountManager.StoreEvent(ctx, userID, record.ID, accountID, activity.DNSRecordCreated, meta) - go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationCreate}) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) return record, nil } @@ -112,6 +119,8 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI var zone *zones.Zone var record *records.Record + var snap *affectedpeers.Snapshot + var change affectedpeers.Change err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID) @@ -141,6 +150,11 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI return fmt.Errorf("failed to update dns record: %w", err) } + change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups} + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + err = transaction.IncrementNetworkSerial(ctx, accountID) if err != nil { return fmt.Errorf("failed to increment network serial: %w", err) @@ -155,7 +169,7 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI meta := record.EventMeta(zone.ID, zone.Name) m.accountManager.StoreEvent(ctx, userID, record.ID, accountID, activity.DNSRecordUpdated, meta) - go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationUpdate}) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) return record, nil } @@ -171,6 +185,8 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI var record *records.Record var zone *zones.Zone + var snap *affectedpeers.Snapshot + var change affectedpeers.Change err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID) @@ -188,6 +204,11 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI return fmt.Errorf("failed to delete dns record: %w", err) } + change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups} + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + err = transaction.IncrementNetworkSerial(ctx, accountID) if err != nil { return fmt.Errorf("failed to increment network serial: %w", err) @@ -202,7 +223,7 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI meta := record.EventMeta(zone.ID, zone.Name) m.accountManager.StoreEvent(ctx, userID, recordID, accountID, activity.DNSRecordDeleted, meta) - go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationDelete}) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) return nil } diff --git a/management/server/affected_peers_zone_test.go b/management/server/affected_peers_zone_test.go new file mode 100644 index 000000000..4d622325c --- /dev/null +++ b/management/server/affected_peers_zone_test.go @@ -0,0 +1,145 @@ +package server + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/controllers/network_map" + "github.com/netbirdio/netbird/management/internals/modules/zones" + "github.com/netbirdio/netbird/management/internals/modules/zones/records" + "github.com/netbirdio/netbird/management/server/affectedpeers" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" +) + +const affectedZoneDomain = "zone.test" + +// createAffectedZone stores a zone distributed to the given groups, optionally with +// one A record so the network map actually ships it. +func createAffectedZone(t *testing.T, s store.Store, accountID, domain string, enabled, withRecord bool, groups []string) *zones.Zone { + t.Helper() + ctx := context.Background() + + zone := zones.NewZone(accountID, domain, domain, enabled, false, groups) + require.NoError(t, s.CreateZone(ctx, zone)) + + if withRecord { + record := records.NewRecord(accountID, zone.ID, "host."+domain, records.RecordTypeA, "10.0.0.1", 300) + require.NoError(t, s.CreateDNSRecord(ctx, record)) + } + + return zone +} + +func TestCollectGroupChange_ZoneLinked(t *testing.T) { + _, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + createAffectedZone(t, s, accountID, affectedZoneDomain, true, true, []string{groupIDs[0]}) + + groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) + assert.Contains(t, groups, groupIDs[0], "group distributed a zone should be affected by its own change") + + groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) + assert.Empty(t, groups, "group not referenced by any zone should not be affected") +} + +func TestCollectGroupChange_UnshippedZoneNotLinked(t *testing.T) { + _, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Disabled zone and zone without records are never shipped by the network map. + createAffectedZone(t, s, accountID, "disabled."+affectedZoneDomain, false, true, []string{groupIDs[0]}) + createAffectedZone(t, s, accountID, "empty."+affectedZoneDomain, true, false, []string{groupIDs[1]}) + + groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0], groupIDs[1]}) + assert.Empty(t, groups, "groups referenced only by unshipped zones should not be affected") +} + +func TestResolveAffectedPeers_ZoneGroupMembershipChange(t *testing.T) { + _, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + + createAffectedZone(t, s, accountID, affectedZoneDomain, true, true, []string{groupIDs[0]}) + + // Same change shape UpdateGroup builds: the group changed as a whole and peer1 + // left it, so peer1 must refresh to drop the zone. + change := affectedpeers.Change{ + ChangedGroupIDs: []string{groupIDs[0]}, + RemovedPeersByGroup: map[string][]string{groupIDs[0]: {peerIDs[1]}}, + } + + result := resolveAffected(t, s, accountID, change) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result, "current and removed members of the zone group should be affected") +} + +func TestResolveAffectedPeers_ZoneDistributionChange(t *testing.T) { + _, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + + // Zone create/update/delete passes old and new distribution groups. + change := affectedpeers.Change{DistributionGroupIDs: []string{groupIDs[0], groupIDs[2]}} + + result := resolveAffected(t, s, accountID, change) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[2]}, result, "only members of the distribution groups should be affected") +} + +// TestAffectedPeers_ZoneGroupUpdate_NewMemberReceivesZone verifies that adding a peer +// to a group referenced only by a zone pushes the zone to the new member and leaves +// unrelated peers alone. +func TestAffectedPeers_ZoneGroupUpdate_NewMemberReceivesZone(t *testing.T) { + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + require.NoError(t, manager.Store.DeletePolicy(ctx, accountID, p.ID)) + } + + zoneGroup := &types.Group{ID: "zone-grp", Name: "ZoneGroup", Peers: []string{peer1.ID}} + require.NoError(t, manager.CreateGroup(ctx, accountID, userID, zoneGroup)) + + createAffectedZone(t, manager.Store, accountID, affectedZoneDomain, true, true, []string{zoneGroup.ID}) + + updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) + updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) + updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + zoneGroup.Peers = []string{peer1.ID, peer2.ID} + require.NoError(t, manager.UpdateGroup(ctx, accountID, userID, zoneGroup)) + + peerShouldReceiveUpdate(t, updMsg1) + msg := receivePeerUpdate(t, updMsg2) + assert.True(t, syncHasCustomZone(msg, affectedZoneDomain+"."), "new zone group member should receive the zone") + peerShouldNotReceiveUpdate(t, updMsg3) +} + +func receivePeerUpdate(t *testing.T, ch <-chan *network_map.UpdateMessage) *network_map.UpdateMessage { + t.Helper() + select { + case msg := <-ch: + require.NotNil(t, msg, "update message should not be nil") + return msg + case <-time.After(peerUpdateTimeout): + require.FailNow(t, "timed out waiting for update message") + return nil + } +} + +func syncHasCustomZone(msg *network_map.UpdateMessage, domain string) bool { + for _, zone := range msg.Update.GetNetworkMap().GetDNSConfig().GetCustomZones() { + if zone.GetDomain() == domain { + return true + } + } + return false +} diff --git a/management/server/affectedpeers/resolver.go b/management/server/affectedpeers/resolver.go index cb2063ac9..895e4fd36 100644 --- a/management/server/affectedpeers/resolver.go +++ b/management/server/affectedpeers/resolver.go @@ -22,6 +22,7 @@ import ( nbdns "github.com/netbirdio/netbird/dns" rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/management/internals/modules/zones" 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" @@ -50,6 +51,7 @@ type Snapshot struct { policies []*types.Policy routes []*route.Route nsGroups []*nbdns.NameServerGroup + zones []*zones.Zone dnsSettings *types.DNSSettings routers []*routerTypes.NetworkRouter resources []*resourceTypes.NetworkResource @@ -127,12 +129,15 @@ func (snap *Snapshot) loadRoutesAndProxy(ctx context.Context, s store.Store, acc return snap.loadProxyServices(ctx, s, accountID) } -// loadDNS loads the nameserver groups and account DNS settings. +// loadDNS loads the nameserver groups, custom DNS zones and account DNS settings. func (snap *Snapshot) loadDNS(ctx context.Context, s store.Store, accountID string) error { var err error if snap.nsGroups, err = s.GetAccountNameServerGroups(ctx, store.LockingStrengthNone, accountID); err != nil { return err } + if snap.zones, err = s.GetAccountZones(ctx, store.LockingStrengthNone, accountID); err != nil { + return err + } snap.dnsSettings, err = s.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID) return err } @@ -357,7 +362,7 @@ func (s policySide) opposite() policySide { // - a changed router/resource/network sits on a NETWORK -> fold the SOURCE side of // the policies whose destination reaches it (and the routers it implies). // -// Routes, nameserver groups, DNS and embedded-proxy services distribute to their own +// Routes, nameserver groups, DNS zones, DNS and embedded-proxy services distribute to their own // member peers, outside the policy graph, and are folded here too. func (r *resolver) walk() { for _, policy := range r.bothSidesPolicies() { @@ -369,6 +374,7 @@ func (r *resolver) walk() { r.collectFromPolicies() r.collectFromRoutes() r.collectFromNameServers() + r.collectFromZones() r.collectFromDNSSettings() r.collectFromNetworkRouters() r.collectFromProxyServices() @@ -829,6 +835,25 @@ func (r *resolver) collectFromNameServers() { } } +// collectFromZones folds the distribution groups of the custom DNS zones that +// reference a linked group. Like nameserver groups, a zone has no opposite side, so +// only a whole-group change folds its groups. Zones the network map does not ship +// (disabled or without records) are skipped. +func (r *resolver) collectFromZones() { + if len(r.linkGroups) == 0 { + return + } + for _, zone := range r.snap.zones { + if !zone.Enabled || len(zone.Records) == 0 { + continue + } + if anyInSet(zone.DistributionGroups, r.linkGroups) { + log.WithContext(r.ctx).Tracef("collectFromZones: zone %s references a linked group -> folding its groups %v (outputGroups only)", zone.ID, zone.DistributionGroups) + r.foldOutputGroups(zone.DistributionGroups) + } + } +} + // collectFromSSHAuthorizedGroups folds the destinations of the enabled SSH rules that // authorize a group whose user membership changed. Those destination peers carry the // group -> user mapping for the groups they authorize, so they refresh even when no