From f7be9c43474f01c8dad80bab7746fd5929128810 Mon Sep 17 00:00:00 2001 From: pascal Date: Thu, 30 Jul 2026 15:45:06 +0200 Subject: [PATCH] legacy network mal and equivalence test --- .../grpc/components_envelope_response.go | 2 +- .../internals/shared/grpc/conversion.go | 8 +- management/internals/shared/grpc/server.go | 2 +- management/server/types/account.go | 48 + management/server/types/account_components.go | 6 +- .../types/legacynmap/account_components.go | 703 +++++++++++ management/server/types/legacynmap/aliases.go | 48 + .../types/legacynmap/component_types.go | 105 ++ .../server/types/legacynmap/converters.go | 128 ++ .../server/types/legacynmap/copied_funcs.go | 284 +++++ management/server/types/legacynmap/doc.go | 7 + .../types/legacynmap/equivalence_test.go | 570 +++++++++ .../types/legacynmap/firewall_helpers.go | 157 +++ .../types/legacynmap/networkmap_components.go | 1034 +++++++++++++++++ .../server/types/legacynmap/proto_legacy.go | 208 ++++ shared/management/types/network.go | 5 + .../management/types/networkmap_components.go | 12 +- 17 files changed, 3318 insertions(+), 9 deletions(-) create mode 100644 management/server/types/legacynmap/account_components.go create mode 100644 management/server/types/legacynmap/aliases.go create mode 100644 management/server/types/legacynmap/component_types.go create mode 100644 management/server/types/legacynmap/converters.go create mode 100644 management/server/types/legacynmap/copied_funcs.go create mode 100644 management/server/types/legacynmap/doc.go create mode 100644 management/server/types/legacynmap/equivalence_test.go create mode 100644 management/server/types/legacynmap/firewall_helpers.go create mode 100644 management/server/types/legacynmap/networkmap_components.go create mode 100644 management/server/types/legacynmap/proto_legacy.go diff --git a/management/internals/shared/grpc/components_envelope_response.go b/management/internals/shared/grpc/components_envelope_response.go index 57c765625..cb6a80ed9 100644 --- a/management/internals/shared/grpc/components_envelope_response.go +++ b/management/internals/shared/grpc/components_envelope_response.go @@ -51,7 +51,7 @@ func ToComponentSyncResponse( // TODO (dmitri) consider using invariants? // enableSSH := computeSSHEnabledForPeer(components, peer) - peerConfig := toPeerConfig(peer, components.Network, dnsName, settings, httpConfig, deviceFlowConfig, enableSSH) + peerConfig := toPeerConfig(peer, components.Network, dnsName, settings, httpConfig, deviceFlowConfig, enableSSH, components.ForceRoutingPeerDNSResolution) includeIPv6 := peer.SupportsIPv6() && peer.IPv6.IsValid() useSourcePrefixes := peer.SupportsSourcePrefixes() diff --git a/management/internals/shared/grpc/conversion.go b/management/internals/shared/grpc/conversion.go index 52b20aafa..debeb8482 100644 --- a/management/internals/shared/grpc/conversion.go +++ b/management/internals/shared/grpc/conversion.go @@ -120,7 +120,7 @@ func toNetbirdConfig(config *nbconfig.Config, turnCredentials *Token, relayToken return nbConfig } -func toPeerConfig(peer *nbpeer.Peer, network *nmdata.Network, dnsName string, settings *types.Settings, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, enableSSH bool) *proto.PeerConfig { +func toPeerConfig(peer *nbpeer.Peer, network *nmdata.Network, dnsName string, settings *types.Settings, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, enableSSH bool, forceRoutingPeerDNS bool) *proto.PeerConfig { netmask, _ := network.Net.Mask.Size() fqdn := peer.FQDN(dnsName) @@ -136,7 +136,7 @@ func toPeerConfig(peer *nbpeer.Peer, network *nmdata.Network, dnsName string, se Address: fmt.Sprintf("%s/%d", peer.IP.String(), netmask), SshConfig: sshConfig, Fqdn: fqdn, - RoutingPeerDnsResolutionEnabled: settings.RoutingPeerDNSResolutionEnabled, + RoutingPeerDnsResolutionEnabled: settings.RoutingPeerDNSResolutionEnabled || peer.ProxyMeta.Embedded || forceRoutingPeerDNS, LazyConnectionEnabled: settings.LazyConnectionEnabled, AutoUpdate: &proto.AutoUpdateSettings{ Version: settings.AutoUpdateVersion, @@ -163,12 +163,12 @@ func ToSyncResponse(ctx context.Context, config *nbconfig.Config, httpConfig *nb useSourcePrefixes := peer.SupportsSourcePrefixes() response := &proto.SyncResponse{ - PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH), + PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH, networkMap.ForceRoutingPeerDNSResolution), NetworkMap: &proto.NetworkMap{ Serial: networkMap.Network.CurrentSerial(), Routes: networkmap.ToProtocolRoutes(networkMap.Routes), DNSConfig: networkmap.ToProtocolDNSConfig(networkMap.DNSConfig, dnsCache, dnsFwdPort), - PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH), + PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH, networkMap.ForceRoutingPeerDNSResolution), }, Checks: toProtocolChecks(ctx, checks), } diff --git a/management/internals/shared/grpc/server.go b/management/internals/shared/grpc/server.go index b3b381e21..a43316ca6 100644 --- a/management/internals/shared/grpc/server.go +++ b/management/internals/shared/grpc/server.go @@ -921,7 +921,7 @@ func (s *Server) prepareLoginResponse(ctx context.Context, peer *nbpeer.Peer, ne // if peer has reached this point then it has logged in loginResp := &proto.LoginResponse{ NetbirdConfig: toNetbirdConfig(s.config, nil, relayToken, nil, settings), - PeerConfig: toPeerConfig(peer, types.TwinNetwork(network), s.networkMapController.GetDNSDomain(settings), settings, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, enableSSH), + PeerConfig: toPeerConfig(peer, types.TwinNetwork(network), s.networkMapController.GetDNSDomain(settings), settings, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, enableSSH, false), Checks: toProtocolChecks(ctx, postureChecks), } diff --git a/management/server/types/account.go b/management/server/types/account.go index 75083e4b1..74ea94865 100644 --- a/management/server/types/account.go +++ b/management/server/types/account.go @@ -1060,6 +1060,54 @@ func (a *Account) GetPeerConnectionResources(ctx context.Context, peer *nbpeer.P return peers, fwRules, authorizedUsers, sshEnabled } +// forcesRoutingPeerDNSResolution reports whether the given peer must run +// routing-peer DNS resolution regardless of the account-global +// RoutingPeerDNSResolutionEnabled setting. It returns true when the peer is a +// router for a domain network resource that is targeted by an enabled +// reverse-proxy service, so the peer's DNS forwarder starts and can resolve +// the target for the embedded proxy peers. Embedded proxy peers themselves are +// handled at PeerConfig build time. +func (a *Account) forcesRoutingPeerDNSResolution(peerID string, routers map[string]map[string]*routerTypes.NetworkRouter) bool { + targeted := a.proxyTargetedDomainResourceIDs() + if len(targeted) == 0 { + return false + } + + for _, resource := range a.NetworkResources { + if resource == nil || !resource.Enabled || resource.Type != resourceTypes.Domain { + continue + } + if _, ok := targeted[resource.ID]; !ok { + continue + } + if _, isRouter := routers[resource.NetworkID][peerID]; isRouter { + return true + } + } + + return false +} + +// proxyTargetedDomainResourceIDs returns the set of domain network resource IDs +// targeted by an enabled, non-terminated reverse-proxy service. +func (a *Account) proxyTargetedDomainResourceIDs() map[string]struct{} { + ids := make(map[string]struct{}) + for _, svc := range a.Services { + if svc == nil || !svc.Enabled || svc.Terminated { + continue + } + for _, target := range svc.Targets { + if target == nil || !target.Enabled { + continue + } + if target.TargetType == service.TargetTypeDomain { + ids[target.TargetId] = struct{}{} + } + } + } + return ids +} + func (a *Account) getAllowedUserIDs() map[string]struct{} { users := make(map[string]struct{}) for _, nbUser := range a.Users { diff --git a/management/server/types/account_components.go b/management/server/types/account_components.go index e482702ba..f42504094 100644 --- a/management/server/types/account_components.go +++ b/management/server/types/account_components.go @@ -104,5 +104,9 @@ func (a *Account) GetPeerNetworkMapComponents( groupIDToUserIDs map[string][]string, ) *NetworkMapComponents { nmd := a.toNetworkMapData(accountZones, validatedPeersMap, resourcePolicies, routers, groupIDToUserIDs) - return nmd.GetPeerNetworkMapComponents(peerID, toTwinCustomZone(peersCustomZone)) + components := nmd.GetPeerNetworkMapComponents(peerID, toTwinCustomZone(peersCustomZone)) + if components != nil { + components.ForceRoutingPeerDNSResolution = a.forcesRoutingPeerDNSResolution(peerID, routers) + } + return components } diff --git a/management/server/types/legacynmap/account_components.go b/management/server/types/legacynmap/account_components.go new file mode 100644 index 000000000..65527f7dd --- /dev/null +++ b/management/server/types/legacynmap/account_components.go @@ -0,0 +1,703 @@ +//go:build nmapequiv + +package legacynmap + +import ( + "context" + "slices" + "time" + + log "github.com/sirupsen/logrus" + + nbdns "github.com/netbirdio/netbird/dns" + "github.com/netbirdio/netbird/management/internals/modules/zones" + routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" + "github.com/netbirdio/netbird/management/server/telemetry" + "github.com/netbirdio/netbird/route" +) + +// GetPeerNetworkMapResult dispatches to either the legacy-NetworkMap path or +// the components path based on the peer's capability and the kill switch. +// Capable peers (PeerCapabilityComponentNetworkMap) get the raw components +// shape — the server skips Calculate() entirely for them, saving CPU +// proportional to the number of capable peers in the account. Legacy peers +// (or any peer when componentsDisabled is true) get the fully-expanded +// NetworkMap as before. + +func GetPeerNetworkMapFromComponents(a *Account, + ctx context.Context, + peerID string, + peersCustomZone nbdns.CustomZone, + accountZones []*zones.Zone, + validatedPeersMap map[string]struct{}, + resourcePolicies map[string][]*Policy, + routers map[string]map[string]*routerTypes.NetworkRouter, + metrics *telemetry.AccountManagerMetrics, + groupIDToUserIDs map[string][]string, +) *NetworkMap { + start := time.Now() + + components := GetPeerNetworkMapComponents(a, + ctx, + peerID, + peersCustomZone, + accountZones, + validatedPeersMap, + resourcePolicies, + routers, + groupIDToUserIDs, + ) + + if components.IsEmpty() { + return &NetworkMap{Network: components.Network} + } + + nm := CalculateNetworkMapFromComponents(ctx, components) + + if metrics != nil { + objectCount := int64(len(nm.Peers) + len(nm.OfflinePeers) + len(nm.Routes) + len(nm.FirewallRules) + len(nm.RoutesFirewallRules)) + metrics.CountNetworkMapObjects(objectCount) + metrics.CountGetPeerNetworkMapDuration(time.Since(start)) + + if objectCount > 5000 { + log.WithContext(ctx).Tracef("account: %s has a total resource count of %d objects from components, "+ + "peers: %d, offline peers: %d, routes: %d, firewall rules: %d, route firewall rules: %d", + a.Id, objectCount, len(nm.Peers), len(nm.OfflinePeers), len(nm.Routes), len(nm.FirewallRules), len(nm.RoutesFirewallRules)) + } + } + + return nm +} + +func GetPeerNetworkMapComponents(a *Account, + ctx context.Context, + peerID string, + peersCustomZone nbdns.CustomZone, + accountZones []*zones.Zone, + validatedPeersMap map[string]struct{}, + resourcePolicies map[string][]*Policy, + routers map[string]map[string]*routerTypes.NetworkRouter, + groupIDToUserIDs map[string][]string, +) *NetworkMapComponents { + peer := a.Peers[peerID] + // this can never happen, things are very wrong if it did + // TODO (dmitri) maybe consider using invariants? + if peer == nil { + log.WithField("peer id", peerID).Error("NetworkMapComponents are computed for a peer missing from the account") + return EmptyNetworkMapComponents(&NetworkMapComponents{ + PeerID: peerID, + Network: a.Network.Copy(), + // must include the target peer as it's required on the client + Peers: map[string]*ComponentPeer{peerID: peerToComponent(peer)}, + }) + } + + if _, ok := validatedPeersMap[peerID]; !ok { + // Mirror legacy graceful-degrade: GetPeerNetworkMapFromComponents + // returns &NetworkMap{Network: a.Network.Copy()} when components is + // nil. Match that floor so the receiving client always sees the + // account Network identifier, not a fully-empty envelope. + return EmptyNetworkMapComponents(&NetworkMapComponents{ + PeerID: peerID, + Network: a.Network.Copy(), + // must include the target peer as it's required on the client + Peers: map[string]*ComponentPeer{peerID: peerToComponent(peer)}, + }) + } + + components := &NetworkMapComponents{ + PeerID: peerID, + Network: a.Network.Copy(), + NameServerGroups: make([]*nbdns.NameServerGroup, 0), + CustomZoneDomain: peersCustomZone.Domain, + ResourcePoliciesMap: make(map[string][]*Policy), + RoutersMap: make(map[string]map[string]*ComponentRouter), + NetworkResources: make([]*ComponentResource, 0), + PostureFailedPeers: make(map[string]map[string]struct{}, len(a.PostureChecks)), + RouterPeers: make(map[string]*ComponentPeer), + NetworkXIDToPublicID: make(map[string]string, len(a.Networks)), + PostureCheckXIDToPublicID: make(map[string]string, len(a.PostureChecks)), + + ForceRoutingPeerDNSResolution: forcesRoutingPeerDNSResolution(a, peerID, routers), + } + for _, n := range a.Networks { + if n != nil { + components.NetworkXIDToPublicID[n.ID] = n.PublicID + } + } + for _, pc := range a.PostureChecks { + if pc != nil { + components.PostureCheckXIDToPublicID[pc.ID] = pc.PublicID + } + } + + components.AccountSettings = &AccountSettingsInfo{ + PeerLoginExpirationEnabled: a.Settings.PeerLoginExpirationEnabled, + PeerLoginExpiration: a.Settings.PeerLoginExpiration, + PeerInactivityExpirationEnabled: a.Settings.PeerInactivityExpirationEnabled, + PeerInactivityExpiration: a.Settings.PeerInactivityExpiration, + } + + components.DNSSettings = &a.DNSSettings + + // relevantPeers always contains the target peer (peerID) + relevantPeers, relevantGroups, relevantPolicies, relevantRoutes, sshReqs := getPeersGroupsPoliciesRoutes(a, ctx, peerID, peer.SSHEnabled, validatedPeersMap, &components.PostureFailedPeers) + + if len(sshReqs.neededGroupIDs) > 0 { + components.GroupIDToUserIDs = filterGroupIDToUserIDs(groupIDToUserIDs, sshReqs.neededGroupIDs) + } + if sshReqs.needAllowedUserIDs { + components.AllowedUserIDs = getAllowedUserIDs(a) + } + + components.Peers = relevantPeers + components.Groups = groupsToComponent(relevantGroups) + components.Policies = relevantPolicies + components.Routes = relevantRoutes + components.AllDNSRecords = filterDNSRecordsByPeers(peersCustomZone.Records, relevantPeers, peer.SupportsIPv6() && peer.IPv6.IsValid()) + + peerGroups := a.GetPeerGroups(peerID) + components.AccountZones = filterPeerAppliedZones(ctx, accountZones, LookupMap(peerGroups)) + components.AccountZones = append(components.AccountZones, a.SynthesizePrivateServiceZones(peerID)...) + + for _, nsGroup := range a.NameServerGroups { + if nsGroup.Enabled { + for _, gID := range nsGroup.Groups { + if _, found := relevantGroups[gID]; found { + components.NameServerGroups = append(components.NameServerGroups, nsGroup) + break + } + } + } + } + + for _, resource := range a.NetworkResources { + if !resource.Enabled { + continue + } + + policies, exists := resourcePolicies[resource.ID] + if !exists { + continue + } + + addSourcePeers := false + + networkRoutingPeers, routerExists := routers[resource.NetworkID] + if routerExists { + if _, ok := networkRoutingPeers[peerID]; ok { + addSourcePeers = true + } + } + + for _, policy := range policies { + if addSourcePeers { + var peers []string + if policy.Rules[0].SourceResource.Type == ResourceTypePeer && policy.Rules[0].SourceResource.ID != "" { + peers = []string{policy.Rules[0].SourceResource.ID} + } else { + peers = getUniquePeerIDsFromGroupsIDs(a, ctx, policy.SourceGroups()) + } + for _, pID := range getPostureValidPeersSaveFailed(a, peers, policy.SourcePostureChecks, validatedPeersMap, &components.PostureFailedPeers) { + if _, exists := components.Peers[pID]; !exists { + components.Peers[pID] = peerToComponent(a.GetPeer(pID)) + } + } + } else { + peerInSources := false + if policy.Rules[0].SourceResource.Type == ResourceTypePeer && policy.Rules[0].SourceResource.ID != "" { + peerInSources = policy.Rules[0].SourceResource.ID == peerID + } else { + for _, groupID := range policy.SourceGroups() { + if group := a.GetGroup(groupID); group != nil && slices.Contains(group.Peers, peerID) { + peerInSources = true + break + } + } + } + if !peerInSources { + continue + } + isValid, pname := validatePostureChecksOnPeerGetFailed(a, ctx, policy.SourcePostureChecks, peerID) + if !isValid && len(pname) > 0 { + if _, ok := components.PostureFailedPeers[pname]; !ok { + components.PostureFailedPeers[pname] = make(map[string]struct{}) + } + components.PostureFailedPeers[pname][peer.ID] = struct{}{} + continue + } + addSourcePeers = true + } + + for _, rule := range policy.Rules { + for _, srcGroupID := range rule.Sources { + if g := a.Groups[srcGroupID]; g != nil { + if _, exists := components.Groups[srcGroupID]; !exists { + components.Groups[srcGroupID] = groupToComponent(g) + } + } + } + for _, dstGroupID := range rule.Destinations { + if g := a.Groups[dstGroupID]; g != nil { + if _, exists := components.Groups[dstGroupID]; !exists { + components.Groups[dstGroupID] = groupToComponent(g) + } + } + } + } + components.ResourcePoliciesMap[resource.ID] = policies + } + + // Only expose router peers and the per-network routers_map when this + // target peer actually has access to the resource (either as a router + // itself or via a policy that includes it as a source). Without this + // gate, every peer's envelope was leaking router peers of every + // network in the account — accounts with many tenants/networks + // shipped tens of unrelated peers in `peers[]` and `routers_map`. + if addSourcePeers { + components.RoutersMap[resource.NetworkID] = routersToComponentMap(networkRoutingPeers) + for peerIDKey := range networkRoutingPeers { + if p := a.Peers[peerIDKey]; p != nil { + cp := components.RouterPeers[peerIDKey] + if cp == nil { + cp = peerToComponent(p) + components.RouterPeers[peerIDKey] = cp + } + if _, exists := components.Peers[peerIDKey]; !exists { + if _, validated := validatedPeersMap[peerIDKey]; validated { + components.Peers[peerIDKey] = cp + } + } + } + } + components.NetworkResources = append(components.NetworkResources, resourceToComponent(resource)) + } + } + + filterGroupPeers(&components.Groups, components.Peers) + filterPostureFailedPeers(&components.PostureFailedPeers, components.Policies, components.ResourcePoliciesMap, components.Peers) + + return components +} + +type sshRequirements struct { + neededGroupIDs map[string]struct{} + needAllowedUserIDs bool +} + +func getPeersGroupsPoliciesRoutes(a *Account, + ctx context.Context, + peerID string, + peerSSHEnabled bool, + validatedPeersMap map[string]struct{}, + postureFailedPeers *map[string]map[string]struct{}, +) (map[string]*ComponentPeer, map[string]*Group, []*Policy, []*route.Route, sshRequirements) { + relevantPeerIDs := make(map[string]*ComponentPeer, len(a.Peers)/4) + relevantGroupIDs := make(map[string]*Group, len(a.Groups)/4) + relevantPolicies := make([]*Policy, 0, len(a.Policies)) + relevantRoutes := make([]*route.Route, 0, len(a.Routes)) + sshReqs := sshRequirements{neededGroupIDs: make(map[string]struct{})} + + relevantPeerIDs[peerID] = peerToComponent(a.GetPeer(peerID)) + + peerGroupSet := make(map[string]struct{}, 8) + for groupID, group := range a.Groups { + if slices.Contains(group.Peers, peerID) { + relevantGroupIDs[groupID] = a.GetGroup(groupID) + peerGroupSet[groupID] = struct{}{} + } + } + + routeAccessControlGroups := make(map[string]struct{}) + for _, r := range a.Routes { + if r == nil { + continue + } + relevant := r.Peer == peerID + if !relevant { + for _, groupID := range r.PeerGroups { + if _, ok := peerGroupSet[groupID]; ok { + relevant = true + break + } + } + } + if !relevant && r.Enabled { + for _, groupID := range r.Groups { + if _, ok := peerGroupSet[groupID]; ok { + relevant = true + break + } + } + } + if !relevant { + continue + } + + for _, groupID := range r.PeerGroups { + relevantGroupIDs[groupID] = a.GetGroup(groupID) + } + for _, groupID := range r.Groups { + relevantGroupIDs[groupID] = a.GetGroup(groupID) + } + if r.Enabled { + for _, groupID := range r.AccessControlGroups { + relevantGroupIDs[groupID] = a.GetGroup(groupID) + routeAccessControlGroups[groupID] = struct{}{} + } + } + + // Include route advertisers in relevantPeerIDs. The envelope + // encoder writes route.peer_index by looking up r.Peer in the + // shipped peers list; if the advertiser is policy-isolated from + // the target peer (no rule edge between them), it would otherwise + // be omitted and the decoder would fail to resolve r.Peer, leaving + // the client without a WG tunnel target for this route. Legacy + // NetworkMap.Routes shipped the WG public key inline, so the + // equivalence path doesn't surface this — but the dependency is + // real once a client actually tries to use the route. + // Gate by validatedPeersMap so non-validated advertisers stay out + // (matches the network-resource router behaviour at the bottom of + // this loop, and the legacy invariant that only validated peers + // reach a client's view). + if r.Peer != "" { + if _, ok := validatedPeersMap[r.Peer]; ok { + if p := a.GetPeer(r.Peer); p != nil { + relevantPeerIDs[r.Peer] = peerToComponent(p) + } + } + } + for _, groupID := range r.PeerGroups { + g := a.GetGroup(groupID) + if g == nil { + continue + } + for _, pid := range g.Peers { + if _, exists := relevantPeerIDs[pid]; exists { + continue + } + if _, ok := validatedPeersMap[pid]; !ok { + continue + } + if p := a.GetPeer(pid); p != nil { + relevantPeerIDs[pid] = peerToComponent(p) + } + } + } + relevantRoutes = append(relevantRoutes, r) + } + + for _, policy := range a.Policies { + if !policy.Enabled { + continue + } + + policyRelevant := false + for _, rule := range policy.Rules { + if !rule.Enabled { + continue + } + + if len(routeAccessControlGroups) > 0 { + for _, destGroupID := range rule.Destinations { + if _, needed := routeAccessControlGroups[destGroupID]; needed { + policyRelevant = true + for _, srcGroupID := range rule.Sources { + relevantGroupIDs[srcGroupID] = a.GetGroup(srcGroupID) + } + for _, dstGroupID := range rule.Destinations { + relevantGroupIDs[dstGroupID] = a.GetGroup(dstGroupID) + } + break + } + } + } + + var sourcePeers, destinationPeers []string + var peerInSources, peerInDestinations bool + + if rule.SourceResource.Type == ResourceTypePeer && rule.SourceResource.ID != "" { + sourcePeers = []string{rule.SourceResource.ID} + if rule.SourceResource.ID == peerID { + peerInSources = true + } + } else { + sourcePeers, peerInSources = getPeersFromGroups(a, ctx, rule.Sources, peerID, policy.SourcePostureChecks, validatedPeersMap, postureFailedPeers) + } + + if rule.DestinationResource.Type == ResourceTypePeer && rule.DestinationResource.ID != "" { + destinationPeers = []string{rule.DestinationResource.ID} + if rule.DestinationResource.ID == peerID { + peerInDestinations = true + } + } else { + destinationPeers, peerInDestinations = getPeersFromGroups(a, ctx, rule.Destinations, peerID, nil, validatedPeersMap, postureFailedPeers) + } + + if peerInSources { + policyRelevant = true + for _, pid := range destinationPeers { + if _, exists := relevantPeerIDs[pid]; !exists { + relevantPeerIDs[pid] = peerToComponent(a.GetPeer(pid)) + } + } + for _, dstGroupID := range rule.Destinations { + relevantGroupIDs[dstGroupID] = a.GetGroup(dstGroupID) + } + } + + if peerInDestinations { + policyRelevant = true + for _, pid := range sourcePeers { + if _, exists := relevantPeerIDs[pid]; !exists { + relevantPeerIDs[pid] = peerToComponent(a.GetPeer(pid)) + } + } + for _, srcGroupID := range rule.Sources { + relevantGroupIDs[srcGroupID] = a.GetGroup(srcGroupID) + } + + if rule.Protocol == PolicyRuleProtocolNetbirdSSH { + switch { + case len(rule.AuthorizedGroups) > 0: + for groupID := range rule.AuthorizedGroups { + sshReqs.neededGroupIDs[groupID] = struct{}{} + } + case rule.AuthorizedUser != "": + default: + sshReqs.needAllowedUserIDs = true + } + } else if PolicyRuleImpliesLegacySSH(rule) && peerSSHEnabled { + sshReqs.needAllowedUserIDs = true + } + } + } + if policyRelevant { + relevantPolicies = append(relevantPolicies, policy) + } + } + + return relevantPeerIDs, relevantGroupIDs, relevantPolicies, relevantRoutes, sshReqs +} + +func getPeersFromGroups(a *Account, ctx context.Context, groups []string, peerID string, sourcePostureChecksIDs []string, + validatedPeersMap map[string]struct{}, postureFailedPeers *map[string]map[string]struct{}) ([]string, bool) { + peerInGroups := false + filteredPeerIDs := make([]string, 0, len(groups)) + seenPeerIds := make(map[string]struct{}, len(groups)) + + for _, gid := range groups { + group := a.GetGroup(gid) + if group == nil { + continue + } + + if group.IsGroupAll() || len(groups) == 1 { + filteredPeerIDs = make([]string, 0, len(group.Peers)) + peerInGroups = false + for _, pid := range group.Peers { + peer, ok := a.Peers[pid] + if !ok || peer == nil { + continue + } + + if _, ok := validatedPeersMap[peer.ID]; !ok { + continue + } + + isValid, pname := validatePostureChecksOnPeerGetFailed(a, ctx, sourcePostureChecksIDs, peer.ID) + if !isValid && len(pname) > 0 { + if _, ok := (*postureFailedPeers)[pname]; !ok { + (*postureFailedPeers)[pname] = make(map[string]struct{}) + } + (*postureFailedPeers)[pname][peer.ID] = struct{}{} + continue + } + + if peer.ID == peerID { + peerInGroups = true + continue + } + + filteredPeerIDs = append(filteredPeerIDs, peer.ID) + } + return filteredPeerIDs, peerInGroups + } + + for _, pid := range group.Peers { + if _, seen := seenPeerIds[pid]; seen { + continue + } + seenPeerIds[pid] = struct{}{} + peer, ok := a.Peers[pid] + if !ok || peer == nil { + continue + } + + if _, ok := validatedPeersMap[peer.ID]; !ok { + continue + } + + isValid, pname := validatePostureChecksOnPeerGetFailed(a, ctx, sourcePostureChecksIDs, peer.ID) + if !isValid && len(pname) > 0 { + if _, ok := (*postureFailedPeers)[pname]; !ok { + (*postureFailedPeers)[pname] = make(map[string]struct{}) + } + (*postureFailedPeers)[pname][peer.ID] = struct{}{} + continue + } + + if peer.ID == peerID { + peerInGroups = true + continue + } + + filteredPeerIDs = append(filteredPeerIDs, peer.ID) + } + } + + return filteredPeerIDs, peerInGroups +} + +func validatePostureChecksOnPeerGetFailed(a *Account, ctx context.Context, sourcePostureChecksID []string, peerID string) (bool, string) { + peer, ok := a.Peers[peerID] + if !ok || peer == nil { + return false, "" + } + + for _, postureChecksID := range sourcePostureChecksID { + postureChecks := a.GetPostureChecks(postureChecksID) + if postureChecks == nil { + continue + } + + for _, check := range postureChecks.GetChecks() { + isValid, _ := check.Check(ctx, *peer) + if !isValid { + return false, postureChecksID + } + } + } + return true, "" +} + +func getPostureValidPeersSaveFailed(a *Account, inputPeers []string, postureChecksIDs []string, validatedPeersMap map[string]struct{}, postureFailedPeers *map[string]map[string]struct{}) []string { + var dest []string + for _, peerID := range inputPeers { + if _, validated := validatedPeersMap[peerID]; !validated { + continue + } + valid, pname := validatePostureChecksOnPeerGetFailed(a, context.Background(), postureChecksIDs, peerID) + if valid { + dest = append(dest, peerID) + continue + } + if _, ok := (*postureFailedPeers)[pname]; !ok { + (*postureFailedPeers)[pname] = make(map[string]struct{}) + } + (*postureFailedPeers)[pname][peerID] = struct{}{} + } + return dest +} + +// filterGroupPeers trims each group's Peers slice to only those peers that +// also appear in `peers`. Groups whose filtered list is empty are NOT +// deleted from the map — they're kept so the components wire encoder can +// still resolve seq references from routes/policies/access-control groups +// that name them. Calculate() tolerates groups with empty Peers (the inner +// loops simply iterate zero times), so retaining them is behaviourally a +// no-op for the legacy path that consumes the same NetworkMapComponents. +func filterGroupPeers(groups *map[string]*ComponentGroup, peers map[string]*ComponentPeer) { + for groupID, groupInfo := range *groups { + filteredPeers := make([]string, 0, len(groupInfo.Peers)) + for _, pid := range groupInfo.Peers { + if _, exists := peers[pid]; exists { + filteredPeers = append(filteredPeers, pid) + } + } + + if len(filteredPeers) != len(groupInfo.Peers) { + ng := *groupInfo + ng.Peers = filteredPeers + (*groups)[groupID] = &ng + } + } +} + +func filterPostureFailedPeers(postureFailedPeers *map[string]map[string]struct{}, policies []*Policy, resourcePoliciesMap map[string][]*Policy, peers map[string]*ComponentPeer) { + if len(*postureFailedPeers) == 0 { + return + } + + referencedPostureChecks := make(map[string]struct{}) + for _, policy := range policies { + for _, checkID := range policy.SourcePostureChecks { + referencedPostureChecks[checkID] = struct{}{} + } + } + for _, resPolicies := range resourcePoliciesMap { + for _, policy := range resPolicies { + for _, checkID := range policy.SourcePostureChecks { + referencedPostureChecks[checkID] = struct{}{} + } + } + } + + for checkID, failedPeers := range *postureFailedPeers { + if _, referenced := referencedPostureChecks[checkID]; !referenced { + delete(*postureFailedPeers, checkID) + continue + } + for peerID := range failedPeers { + if _, exists := peers[peerID]; !exists { + delete(failedPeers, peerID) + } + } + if len(failedPeers) == 0 { + delete(*postureFailedPeers, checkID) + } + } +} + +func filterDNSRecordsByPeers(records []nbdns.SimpleRecord, peers map[string]*ComponentPeer, includeIPv6 bool) []nbdns.SimpleRecord { + if len(records) == 0 || len(peers) == 0 { + return nil + } + + // Include both v4 and v6 addresses so AAAA records (whose RData is an IPv6 + // address) are not filtered out when peers have IPv6 assigned. When the + // requesting peer doesn't have IPv6, omit v6 IPs so AAAA records get dropped. + peerIPs := make(map[string]struct{}, len(peers)*2) + for _, peer := range peers { + if peer == nil { + continue + } + peerIPs[peer.IP.String()] = struct{}{} + if includeIPv6 && peer.IPv6.IsValid() { + peerIPs[peer.IPv6.String()] = struct{}{} + } + } + + filteredRecords := make([]nbdns.SimpleRecord, 0, len(records)) + for _, record := range records { + if _, exists := peerIPs[record.RData]; exists { + filteredRecords = append(filteredRecords, record) + } + } + + return filteredRecords +} + +func filterGroupIDToUserIDs(fullMap map[string][]string, neededGroupIDs map[string]struct{}) map[string][]string { + if len(neededGroupIDs) == 0 { + return nil + } + + filtered := make(map[string][]string, len(neededGroupIDs)) + for groupID := range neededGroupIDs { + if users, ok := fullMap[groupID]; ok { + filtered[groupID] = users + } + } + return filtered +} diff --git a/management/server/types/legacynmap/aliases.go b/management/server/types/legacynmap/aliases.go new file mode 100644 index 000000000..9e5cceb36 --- /dev/null +++ b/management/server/types/legacynmap/aliases.go @@ -0,0 +1,48 @@ +//go:build nmapequiv + +// Package legacynmap is a frozen copy of main's Account → NetworkMapComponents → +// NetworkMap → proto path, used only by the main-vs-branch equivalence test. +// It is build-tagged so it never compiles into production binaries, and it lives +// in its own package so it cannot reach this tree's unexported helpers — a +// divergence can therefore never be hidden by the two sides sharing code. +// +// Delete this package once the nmdata refactor is validated. +// +// Types below are aliased rather than copied because they are byte-identical +// between main and this branch. Anything that drifted is copied instead; see +// converters.go and copied_funcs.go. +package legacynmap + +import ( + types "github.com/netbirdio/netbird/management/server/types" + sharedtypes "github.com/netbirdio/netbird/shared/management/types" +) + +type ( + Account = types.Account + + DNSSettings = sharedtypes.DNSSettings + FirewallRule = sharedtypes.FirewallRule + ForwardingRule = sharedtypes.ForwardingRule + Group = sharedtypes.Group + Network = sharedtypes.Network + Policy = sharedtypes.Policy + PolicyRule = sharedtypes.PolicyRule + Resource = sharedtypes.Resource + RulePortRange = sharedtypes.RulePortRange + RouteFirewallRule = sharedtypes.RouteFirewallRule +) + +const ( + FirewallRuleDirectionIN = sharedtypes.FirewallRuleDirectionIN + FirewallRuleDirectionOUT = sharedtypes.FirewallRuleDirectionOUT + + PolicyRuleProtocolALL = sharedtypes.PolicyRuleProtocolALL + PolicyRuleProtocolTCP = sharedtypes.PolicyRuleProtocolTCP + PolicyRuleProtocolNetbirdSSH = sharedtypes.PolicyRuleProtocolNetbirdSSH + PolicyTrafficActionAccept = sharedtypes.PolicyTrafficActionAccept + ResourceTypePeer = sharedtypes.ResourceTypePeer + + AllowedIPsFormat = sharedtypes.AllowedIPsFormat + AllowedIPsV6Format = sharedtypes.AllowedIPsV6Format +) diff --git a/management/server/types/legacynmap/component_types.go b/management/server/types/legacynmap/component_types.go new file mode 100644 index 000000000..373b0c04a --- /dev/null +++ b/management/server/types/legacynmap/component_types.go @@ -0,0 +1,105 @@ +//go:build nmapequiv + +package legacynmap + +import ( + "net/netip" + "time" +) + +// ComponentPeer is the self-contained peer representation used by +// NetworkMapComponents and the calculated NetworkMap. It carries exactly the +// subset of peer data that crosses the components wire format, so the shared +// calculation layer stays independent of the management server's domain +// types. +type ComponentPeer struct { + ID string + Key string + IP netip.Addr + IPv6 netip.Addr + DNSLabel string + SSHKey string + SSHEnabled bool + ServerSSHAllowed bool + AgentVersion string + SupportsSourcePrefixes bool + SupportsIPv6 bool + LoginExpirationEnabled bool + AddedWithSSOLogin bool + LastLogin time.Time +} + +// FQDN returns the peer's FQDN combined of the peer's DNS label and the system's DNS domain. +func (p *ComponentPeer) FQDN(dnsDomain string) string { + if dnsDomain == "" { + return "" + } + return p.DNSLabel + "." + dnsDomain +} + +// LoginExpired indicates whether the peer's login has expired, mirroring the +// server-side peer semantics: only SSO-added peers with login expiration +// enabled can expire. +func (p *ComponentPeer) LoginExpired(expiresIn time.Duration) (bool, time.Duration) { + if !p.AddedWithSSOLogin || !p.LoginExpirationEnabled { + return false, 0 + } + timeLeft := time.Until(p.LastLogin.Add(expiresIn)) + return timeLeft <= 0, timeLeft +} + +// GroupAllName is the reserved name of the default group that contains every peer in an account. +const GroupAllName = "All" + +// ComponentGroup is the self-contained group representation used by +// NetworkMapComponents: just the membership view the network-map calculation +// needs, without the server's storage fields. +type ComponentGroup struct { + ID string + PublicID string + Name string + Peers []string +} + +// IsGroupAll checks if the group is a default "All" group. +func (g *ComponentGroup) IsGroupAll() bool { + return g.Name == GroupAllName +} + +// ComponentRouter is the self-contained network-router representation used by +// NetworkMapComponents. +type ComponentRouter struct { + NetworkID string + PublicID string + Peer string + PeerGroups []string + Masquerade bool + Metric int + Enabled bool +} + +// ComponentResourceType mirrors the network-resource type enum on the +// components wire format. +type ComponentResourceType string + +const ( + ComponentResourceHost ComponentResourceType = "host" + ComponentResourceSubnet ComponentResourceType = "subnet" + ComponentResourceDomain ComponentResourceType = "domain" +) + +// ComponentResource is the self-contained network-resource representation +// used by NetworkMapComponents. +type ComponentResource struct { + ID string + PublicID string + NetworkID string + AccountID string + Name string + Description string + Type ComponentResourceType + Address string + Domain string + Prefix netip.Prefix + Enabled bool +} diff --git a/management/server/types/legacynmap/converters.go b/management/server/types/legacynmap/converters.go new file mode 100644 index 000000000..9578ccdda --- /dev/null +++ b/management/server/types/legacynmap/converters.go @@ -0,0 +1,128 @@ +//go:build nmapequiv + +package legacynmap + +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" + nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/route" +) + +// NetworkMap is main's shape. It is copied rather than aliased because this +// branch's NetworkMap dropped ForceRoutingPeerDNSResolution, which main threads +// into PeerConfig.RoutingPeerDnsResolutionEnabled. +type NetworkMap struct { + Peers []*ComponentPeer + Network *Network + Routes []*route.Route + DNSConfig nbdns.Config + OfflinePeers []*ComponentPeer + FirewallRules []*FirewallRule + RoutesFirewallRules []*RouteFirewallRule + ForwardingRules []*ForwardingRule + AuthorizedUsers map[string]map[string]struct{} + EnableSSH bool + // ForceRoutingPeerDNSResolution forces the peer to run/use routing-peer DNS + // resolution regardless of the account-global setting, for reverse-proxy + // domain targets. + ForceRoutingPeerDNSResolution bool +} + +// The ToComponent converters below are main's methods, re-expressed as free +// functions because their receivers live in packages this one cannot extend. +// Bodies are otherwise unchanged. + +func peerToComponent(p *nbpeer.Peer) *ComponentPeer { + if p == nil { + return nil + } + cp := &ComponentPeer{ + ID: p.ID, + Key: p.Key, + IP: p.IP, + IPv6: p.IPv6, + DNSLabel: p.DNSLabel, + SSHKey: p.SSHKey, + SSHEnabled: p.SSHEnabled, + ServerSSHAllowed: p.Meta.Flags.ServerSSHAllowed, + AgentVersion: p.Meta.WtVersion, + SupportsSourcePrefixes: p.SupportsSourcePrefixes(), + SupportsIPv6: p.SupportsIPv6(), + LoginExpirationEnabled: p.LoginExpirationEnabled, + AddedWithSSOLogin: p.AddedWithSSOLogin(), + } + if p.LastLogin != nil { + cp.LastLogin = *p.LastLogin + } + return cp +} + +func groupToComponent(g *Group) *ComponentGroup { + if g == nil { + return nil + } + return &ComponentGroup{ + ID: g.ID, + PublicID: g.PublicID, + Name: g.Name, + Peers: g.Peers, + } +} + +func groupsToComponent(groups map[string]*Group) map[string]*ComponentGroup { + if groups == nil { + return nil + } + out := make(map[string]*ComponentGroup, len(groups)) + for id, g := range groups { + out[id] = groupToComponent(g) + } + return out +} + +func routerToComponent(n *routerTypes.NetworkRouter) *ComponentRouter { + if n == nil { + return nil + } + return &ComponentRouter{ + NetworkID: n.NetworkID, + PublicID: n.PublicID, + Peer: n.Peer, + PeerGroups: n.PeerGroups, + Masquerade: n.Masquerade, + Metric: n.Metric, + Enabled: n.Enabled, + } +} + +func routersToComponentMap(routers map[string]*routerTypes.NetworkRouter) map[string]*ComponentRouter { + if routers == nil { + return nil + } + out := make(map[string]*ComponentRouter, len(routers)) + for id, r := range routers { + out[id] = routerToComponent(r) + } + return out +} + +func resourceToComponent(n *resourceTypes.NetworkResource) *ComponentResource { + if n == nil { + return nil + } + return &ComponentResource{ + ID: n.ID, + PublicID: n.PublicID, + NetworkID: n.NetworkID, + AccountID: n.AccountID, + Name: n.Name, + Description: n.Description, + Type: ComponentResourceType(n.Type), + Address: n.Address, + Domain: n.Domain, + Prefix: n.Prefix, + Enabled: n.Enabled, + } +} diff --git a/management/server/types/legacynmap/copied_funcs.go b/management/server/types/legacynmap/copied_funcs.go new file mode 100644 index 000000000..8f09b4b5e --- /dev/null +++ b/management/server/types/legacynmap/copied_funcs.go @@ -0,0 +1,284 @@ +//go:build nmapequiv + +package legacynmap + +import ( + "context" + "fmt" + "strconv" + "strings" + + "github.com/miekg/dns" + log "github.com/sirupsen/logrus" + + nbdns "github.com/netbirdio/netbird/dns" + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/management/internals/modules/zones" + "github.com/netbirdio/netbird/management/internals/modules/zones/records" + resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" + routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" + nbroute "github.com/netbirdio/netbird/route" +) + +func GenerateRouteFirewallRules(ctx context.Context, route *nbroute.Route, rule *PolicyRule, groupPeers []*ComponentPeer, direction int, includeIPv6 bool) []*RouteFirewallRule { + rulesExists := make(map[string]struct{}) + rules := make([]*RouteFirewallRule, 0) + + v4Sources, v6Sources := splitPeerSourcesByFamily(groupPeers) + + isV6Route := route.Network.Addr().Is6() + + // Skip v6 destination routes entirely for peers without IPv6 support + if isV6Route && !includeIPv6 { + return rules + } + + // Pick sources matching the destination family + sourceRanges := v4Sources + if isV6Route { + sourceRanges = v6Sources + } + + baseRule := RouteFirewallRule{ + PolicyID: rule.PolicyID, + RouteID: route.ID, + SourceRanges: sourceRanges, + Action: string(rule.Action), + Destination: route.Network.String(), + Protocol: string(rule.Protocol), + Domains: route.Domains, + IsDynamic: route.IsDynamic(), + } + + if len(rule.Ports) == 0 { + rules = append(rules, generateRulesWithPortRanges(baseRule, rule, rulesExists)...) + } else { + rules = append(rules, generateRulesWithPorts(ctx, baseRule, rule, rulesExists)...) + } + + // Generate v6 counterpart for dynamic routes and 0.0.0.0/0 exit node routes. + isDefaultV4 := !isV6Route && route.Network.Bits() == 0 + if includeIPv6 && (route.IsDynamic() || isDefaultV4) && len(v6Sources) > 0 { + v6Rule := baseRule + v6Rule.SourceRanges = v6Sources + if isDefaultV4 { + v6Rule.Destination = "::/0" + v6Rule.RouteID = route.ID + "-v6-default" + } + if len(rule.Ports) == 0 { + rules = append(rules, generateRulesWithPortRanges(v6Rule, rule, rulesExists)...) + } else { + rules = append(rules, generateRulesWithPorts(ctx, v6Rule, rule, rulesExists)...) + } + } + + return rules +} + +func filterPeerAppliedZones(ctx context.Context, accountZones []*zones.Zone, peerGroups LookupMap) []nbdns.CustomZone { + var customZones []nbdns.CustomZone + + if len(peerGroups) == 0 { + return customZones + } + + for _, zone := range accountZones { + if !zone.Enabled || len(zone.Records) == 0 { + continue + } + + hasAccess := false + for _, distGroupID := range zone.DistributionGroups { + if _, found := peerGroups[distGroupID]; found { + hasAccess = true + break + } + } + + if !hasAccess { + continue + } + + simpleRecords := make([]nbdns.SimpleRecord, 0, len(zone.Records)) + for _, record := range zone.Records { + var recordType int + rData := record.Content + + switch record.Type { + case records.RecordTypeA: + recordType = int(dns.TypeA) + case records.RecordTypeAAAA: + recordType = int(dns.TypeAAAA) + case records.RecordTypeCNAME: + recordType = int(dns.TypeCNAME) + rData = dns.Fqdn(record.Content) + default: + log.WithContext(ctx).Warnf("unknown DNS record type %s for record %s", record.Type, record.ID) + continue + } + + simpleRecords = append(simpleRecords, nbdns.SimpleRecord{ + Name: dns.Fqdn(record.Name), + Type: recordType, + Class: nbdns.DefaultClass, + TTL: record.TTL, + RData: rData, + }) + } + + customZones = append(customZones, nbdns.CustomZone{ + Domain: dns.Fqdn(zone.Domain), + Records: simpleRecords, + SearchDomainDisabled: !zone.EnableSearchDomain, + NonAuthoritative: true, + }) + } + + return customZones +} + +func getAllowedUserIDs(a *Account) map[string]struct{} { + users := make(map[string]struct{}) + for _, nbUser := range a.Users { + if !nbUser.IsBlocked() && !nbUser.IsServiceUser { + users[nbUser.Id] = struct{}{} + } + } + return users +} + +func getUniquePeerIDsFromGroupsIDs(a *Account, ctx context.Context, groups []string) []string { + peerIDs := make(map[string]struct{}, len(groups)) // we expect at least one peer per group as initial capacity + for _, groupID := range groups { + group := a.GetGroup(groupID) + if group == nil { + log.WithContext(ctx).Warnf("group %s doesn't exist under account %s, will continue map generation without it", groupID, a.Id) + continue + } + + if group.IsGroupAll() || len(groups) == 1 { + return group.Peers + } + + for _, peerID := range group.Peers { + peerIDs[peerID] = struct{}{} + } + } + + ids := make([]string, 0, len(peerIDs)) + for peerID := range peerIDs { + ids = append(ids, peerID) + } + + return ids +} + +func forcesRoutingPeerDNSResolution(a *Account, peerID string, routers map[string]map[string]*routerTypes.NetworkRouter) bool { + targeted := proxyTargetedDomainResourceIDs(a) + if len(targeted) == 0 { + return false + } + + for _, resource := range a.NetworkResources { + if resource == nil || !resource.Enabled || resource.Type != resourceTypes.Domain { + continue + } + if _, ok := targeted[resource.ID]; !ok { + continue + } + if _, isRouter := routers[resource.NetworkID][peerID]; isRouter { + return true + } + } + + return false +} + +func proxyTargetedDomainResourceIDs(a *Account) map[string]struct{} { + ids := make(map[string]struct{}) + for _, svc := range a.Services { + if svc == nil || !svc.Enabled || svc.Terminated { + continue + } + for _, target := range svc.Targets { + if target == nil || !target.Enabled { + continue + } + if target.TargetType == service.TargetTypeDomain { + ids[target.TargetId] = struct{}{} + } + } + } + return ids +} + +func splitPeerSourcesByFamily(groupPeers []*ComponentPeer) (v4, v6 []string) { + v4 = make([]string, 0, len(groupPeers)) + v6 = make([]string, 0, len(groupPeers)) + for _, peer := range groupPeers { + if peer == nil { + continue + } + v4 = append(v4, fmt.Sprintf(AllowedIPsFormat, peer.IP)) + if peer.IPv6.IsValid() { + v6 = append(v6, fmt.Sprintf(AllowedIPsV6Format, peer.IPv6)) + } + } + return +} + +func generateRulesWithPortRanges(baseRule RouteFirewallRule, rule *PolicyRule, rulesExists map[string]struct{}) []*RouteFirewallRule { + rules := make([]*RouteFirewallRule, 0) + + ruleIDBase := generateRuleIDBase(rule, baseRule) + if len(rule.Ports) == 0 { + if len(rule.PortRanges) == 0 { + if _, ok := rulesExists[ruleIDBase]; !ok { + rulesExists[ruleIDBase] = struct{}{} + rules = append(rules, &baseRule) + } + } else { + for _, portRange := range rule.PortRanges { + ruleID := fmt.Sprintf("%s%d-%d", ruleIDBase, portRange.Start, portRange.End) + if _, ok := rulesExists[ruleID]; !ok { + rulesExists[ruleID] = struct{}{} + pr := baseRule + pr.PortRange = portRange + rules = append(rules, &pr) + } + } + } + return rules + } + + return rules +} + +func generateRulesWithPorts(ctx context.Context, baseRule RouteFirewallRule, rule *PolicyRule, rulesExists map[string]struct{}) []*RouteFirewallRule { + rules := make([]*RouteFirewallRule, 0) + ruleIDBase := generateRuleIDBase(rule, baseRule) + + for _, port := range rule.Ports { + ruleID := ruleIDBase + port + if _, ok := rulesExists[ruleID]; ok { + continue + } + rulesExists[ruleID] = struct{}{} + + pr := baseRule + p, err := strconv.ParseUint(port, 10, 16) + if err != nil { + log.WithContext(ctx).Errorf("failed to parse port %s for rule: %s", port, rule.ID) + continue + } + + pr.Port = uint16(p) + rules = append(rules, &pr) + } + + return rules +} + +func generateRuleIDBase(rule *PolicyRule, baseRule RouteFirewallRule) string { + return rule.ID + strings.Join(baseRule.SourceRanges, ",") + strconv.Itoa(FirewallRuleDirectionIN) + baseRule.Protocol + baseRule.Action +} diff --git a/management/server/types/legacynmap/doc.go b/management/server/types/legacynmap/doc.go new file mode 100644 index 000000000..22661d404 --- /dev/null +++ b/management/server/types/legacynmap/doc.go @@ -0,0 +1,7 @@ +// Package legacynmap holds a frozen copy of main's network-map computation, +// used only by the main-vs-branch proto equivalence test. All real content is +// behind the nmapequiv build tag; this file exists so the package is still valid +// for untagged builds and `go test ./...`. +// +// Delete this package once the nmdata refactor is validated. +package legacynmap diff --git a/management/server/types/legacynmap/equivalence_test.go b/management/server/types/legacynmap/equivalence_test.go new file mode 100644 index 000000000..9b6eab651 --- /dev/null +++ b/management/server/types/legacynmap/equivalence_test.go @@ -0,0 +1,570 @@ +//go:build nmapequiv + +// Main-vs-branch equivalence check. For every peer of every account in a real +// Postgres copy it computes the client-facing proto.NetworkMap twice: +// +// - legacy path: main's Account → NetworkMapComponents → Calculate → proto +// (the frozen copy in this package) +// - new path: this branch's Account → NetworkMapData → components → +// Calculate → ToSyncResponse → proto +// +// proto.NetworkMap is generated code identical in both trees, which is what +// makes it the one usable comparison surface — the intermediate Go types differ +// by design. proto.Equal would trip over repeated-field ordering, so both sides +// are canonicalized first. +// +// NETBIRD_STORE_ENGINE_POSTGRES_DSN='...' go test -tags nmapequiv \ +// -run TestNetworkMapProtoEquivalence -count=1 -timeout 60m \ +// ./management/server/types/legacynmap/ +// +// Accounts are loaded one at a time and released between iterations, so peak +// memory tracks the largest single account rather than the whole database. +// +// Env knobs: NETMAP_ACCOUNTS (comma-separated ids, skips discovery), +// NETMAP_MAX_ACCOUNTS (0 = all), NETMAP_MAX_PEERS (0 = all). Fails at the +// first divergence. +package legacynmap_test + +import ( + "bytes" + "cmp" + "context" + "os" + "runtime" + "runtime/debug" + "slices" + "sort" + "strconv" + "strings" + "testing" + + "github.com/stretchr/testify/require" + "google.golang.org/protobuf/encoding/prototext" + goproto "google.golang.org/protobuf/proto" + "gorm.io/driver/postgres" + "gorm.io/gorm" + gormlogger "gorm.io/gorm/logger" + + nbdns "github.com/netbirdio/netbird/dns" + "github.com/netbirdio/netbird/management/internals/controllers/network_map/controller/cache" + mgmtgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/management/server/types/legacynmap" + "github.com/netbirdio/netbird/shared/management/proto" +) + +const ( + equivDNSName = "netbird.cloud" + progressEvery = 5000 +) + +type equivStats struct { + accounts int + peersChecked int + skippedNilNM int +} + +func TestNetworkMapProtoEquivalence(t *testing.T) { + if testing.Short() { + t.Skip("prod-db equivalence test, skipped in short mode") + } + dsn := equivDSN() + if dsn == "" { + t.Skip("NETBIRD_STORE_ENGINE_POSTGRES_DSN not set") + } + + ctx := context.Background() + // skipMigration=true: this reads a restored production copy and must not + // alter its schema. Flip to false only if reads fail on an older dump. + testStore, err := store.NewPostgresqlStore(ctx, dsn, nil, true) + require.NoError(t, err, "connect to postgres") + t.Cleanup(func() { testStore.Close(ctx) }) + + accountIDs := equivAccountIDs(t, dsn) + require.NotEmpty(t, accountIDs, "no accounts selected") + + stats := &equivStats{accounts: len(accountIDs)} + maxPeers := envInt("NETMAP_MAX_PEERS", 0) + + for i, accountID := range accountIDs { + account, err := testStore.GetAccount(ctx, accountID) + if err != nil { + t.Logf("account %s: load failed, skipping: %v", accountID, err) + continue + } + + checkAccount(ctx, t, account, maxPeers, stats) + + account = nil + debug.FreeOSMemory() + + if i%progressEvery == 0 { + var ms runtime.MemStats + runtime.ReadMemStats(&ms) + t.Logf("progress: accounts=%d/%d peers_checked=%d heap=%dMiB", i, len(accountIDs), stats.peersChecked, ms.HeapAlloc>>20) + } + } + + t.Logf("equivalence: accounts=%d peers_checked=%d skipped_nil_nm=%d — no divergence", + stats.accounts, stats.peersChecked, stats.skippedNilNM) +} + +// checkAccount compares both paths for every peer of one account. Nothing is +// retained across peers, so memory stays flat within an account. +func checkAccount(ctx context.Context, t *testing.T, account *types.Account, maxPeers int, stats *equivStats) { + t.Helper() + + if len(account.Peers) == 0 { + return + } + + validated := make(map[string]struct{}, len(account.Peers)) + peerIDs := make([]string, 0, len(account.Peers)) + for peerID := range account.Peers { + validated[peerID] = struct{}{} + peerIDs = append(peerIDs, peerID) + } + sort.Strings(peerIDs) + if maxPeers > 0 && len(peerIDs) > maxPeers { + peerIDs = peerIDs[:maxPeers] + } + + resourcePolicies := account.GetResourcePoliciesMap() + routers := account.GetResourceRoutersMap() + groupUsers := account.GetActiveGroupUsers() + + settings := account.Settings + if settings == nil { + settings = &types.Settings{} + } + + for _, peerID := range peerIDs { + peer := account.Peers[peerID] + if peer == nil { + continue + } + + // NEW PATH — this branch, through the production conversion. + newNM := account.GetPeerNetworkMapFromComponents( + ctx, peerID, nbdns.CustomZone{}, nil, validated, resourcePolicies, routers, nil, groupUsers, + ) + if newNM == nil { + stats.skippedNilNM++ + continue + } + newProto := mgmtgrpc.ToSyncResponse( + ctx, nil, nil, nil, peer, nil, nil, newNM, equivDNSName, nil, + &cache.DNSConfigCache{}, settings, settings.Extra, nil, 0, + ).NetworkMap + + // LEGACY PATH — main's frozen copy. + legacyNM := legacynmap.GetPeerNetworkMapFromComponents( + account, ctx, peerID, nbdns.CustomZone{}, nil, validated, resourcePolicies, routers, nil, groupUsers, + ) + if legacyNM == nil { + t.Fatalf("after %d peers: account=%s peer=%s legacy NetworkMap nil, new non-nil", stats.peersChecked, account.Id, peerID) + } + // A separate cache per side: sharing one would let the first path + // populate entries the second then reuses, which can mask a real diff. + legacyProto := legacynmap.ToProtoNetworkMap( + ctx, peer, legacyNM, equivDNSName, settings, nil, &cache.DNSConfigCache{}, 0, + ) + + canonicalize(legacyProto) + canonicalize(newProto) + stats.peersChecked++ + + if !goproto.Equal(legacyProto, newProto) { + t.Fatalf("after %d peers: %s", stats.peersChecked, describeDivergence(legacyProto, newProto, account.Id, peerID)) + } + } +} + +func equivDSN() string { + if dsn := os.Getenv("NETBIRD_STORE_ENGINE_POSTGRES_DSN"); dsn != "" { + return dsn + } + return os.Getenv("NB_STORE_ENGINE_POSTGRES_DSN") +} + +// equivAccountIDs lists account ids with an id-only query. store.GetAllAccounts +// would hydrate every account in the database before the first comparison runs. +// Sorting happens in Go so the order does not depend on database collation. +func equivAccountIDs(t *testing.T, dsn string) []string { + t.Helper() + + if ids := strings.TrimSpace(os.Getenv("NETMAP_ACCOUNTS")); ids != "" { + var out []string + for _, id := range strings.Split(ids, ",") { + if id = strings.TrimSpace(id); id != "" { + out = append(out, id) + } + } + sort.Strings(out) + return out + } + + db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{Logger: gormlogger.Discard}) + require.NoError(t, err, "open id-listing connection") + defer func() { + if sqlDB, err := db.DB(); err == nil { + sqlDB.Close() + } + }() + + var ids []string + require.NoError(t, db.Model(&types.Account{}).Pluck("id", &ids).Error) + sort.Strings(ids) + + if max := envInt("NETMAP_MAX_ACCOUNTS", 0); max > 0 && len(ids) > max { + ids = ids[:max] + } + return ids +} + +func envInt(name string, def int) int { + if v := os.Getenv(name); v != "" { + if n, err := strconv.Atoi(v); err == nil { + return n + } + } + return def +} + +// canonicalize sorts every repeated field by a stable key. Both paths iterate Go +// maps while building these slices, so order can differ even when the content is +// identical; without this proto.Equal reports noise. +func canonicalize(nm *proto.NetworkMap) { + if nm == nil { + return + } + slices.SortFunc(nm.RemotePeers, cmpRemotePeer) + slices.SortFunc(nm.OfflinePeers, cmpRemotePeer) + slices.SortFunc(nm.Routes, cmpRoute) + slices.SortFunc(nm.FirewallRules, cmpFirewallRule) + slices.SortFunc(nm.RoutesFirewallRules, cmpRouteFirewallRule) + slices.SortFunc(nm.ForwardingRules, cmpForwardingRule) + + for _, r := range nm.FirewallRules { + slices.SortFunc(r.SourcePrefixes, bytes.Compare) + } + for _, r := range nm.RoutesFirewallRules { + slices.Sort(r.SourceRanges) + } + canonicalizeDNSConfig(nm.DNSConfig) + canonicalizeSSHAuth(nm.SshAuth) +} + +func canonicalizeDNSConfig(d *proto.DNSConfig) { + if d == nil { + return + } + for _, g := range d.NameServerGroups { + if g == nil { + continue + } + slices.Sort(g.Domains) + slices.SortFunc(g.NameServers, func(a, b *proto.NameServer) int { + if a == nil || b == nil { + return boolCmp(a == nil, b == nil) + } + if c := cmp.Compare(a.IP, b.IP); c != 0 { + return c + } + if c := cmp.Compare(a.Port, b.Port); c != 0 { + return c + } + return cmp.Compare(a.NSType, b.NSType) + }) + } + slices.SortFunc(d.NameServerGroups, func(a, b *proto.NameServerGroup) int { + return cmp.Compare(nsgKey(a), nsgKey(b)) + }) + for _, z := range d.CustomZones { + if z == nil { + continue + } + slices.SortFunc(z.Records, cmpSimpleRecord) + } + slices.SortFunc(d.CustomZones, func(a, b *proto.CustomZone) int { + if a == nil || b == nil { + return boolCmp(a == nil, b == nil) + } + return cmp.Compare(a.Domain, b.Domain) + }) +} + +// canonicalizeSSHAuth sorts AuthorizedUsers and re-keys MachineUsers.Indexes +// against the new ordering, preserving which machine user maps to which hashes. +func canonicalizeSSHAuth(s *proto.SSHAuth) { + if s == nil || len(s.AuthorizedUsers) == 0 { + return + } + type hashed struct { + bytes []byte + old uint32 + } + entries := make([]hashed, len(s.AuthorizedUsers)) + for i, b := range s.AuthorizedUsers { + entries[i] = hashed{bytes: b, old: uint32(i)} + } + slices.SortFunc(entries, func(a, b hashed) int { return bytes.Compare(a.bytes, b.bytes) }) + + remap := make(map[uint32]uint32, len(entries)) + sorted := make([][]byte, len(entries)) + for newIdx, e := range entries { + remap[e.old] = uint32(newIdx) + sorted[newIdx] = e.bytes + } + s.AuthorizedUsers = sorted + + for _, mu := range s.MachineUsers { + if mu == nil { + continue + } + for i, oldIdx := range mu.Indexes { + if newIdx, ok := remap[oldIdx]; ok { + mu.Indexes[i] = newIdx + } + } + slices.Sort(mu.Indexes) + } +} + +func boolCmp(a, b bool) int { + if a == b { + return 0 + } + if a { + return 1 + } + return -1 +} + +func nsgKey(g *proto.NameServerGroup) string { + if g == nil { + return "" + } + var parts []string + for _, ns := range g.NameServers { + if ns == nil { + continue + } + parts = append(parts, ns.IP+":"+strconv.FormatInt(ns.Port, 10)+":"+strconv.FormatInt(ns.NSType, 10)) + } + slices.Sort(parts) + key := strings.Join(parts, ",") + domains := append([]string(nil), g.Domains...) + slices.Sort(domains) + key += "|" + strings.Join(domains, "|") + if g.Primary { + key += "|P" + } + if g.SearchDomainsEnabled { + key += "|S" + } + return key +} + +func cmpSimpleRecord(a, b *proto.SimpleRecord) int { + if a == nil || b == nil { + return boolCmp(a == nil, b == nil) + } + if c := cmp.Compare(a.Name, b.Name); c != 0 { + return c + } + if c := cmp.Compare(a.Type, b.Type); c != 0 { + return c + } + if c := cmp.Compare(a.Class, b.Class); c != 0 { + return c + } + if c := cmp.Compare(a.RData, b.RData); c != 0 { + return c + } + return cmp.Compare(a.TTL, b.TTL) +} + +func cmpRemotePeer(a, b *proto.RemotePeerConfig) int { + if a == nil || b == nil { + return boolCmp(a == nil, b == nil) + } + return cmp.Compare(a.WgPubKey, b.WgPubKey) +} + +func cmpRoute(a, b *proto.Route) int { + if a == nil || b == nil { + return boolCmp(a == nil, b == nil) + } + if c := cmp.Compare(a.ID, b.ID); c != 0 { + return c + } + if c := cmp.Compare(a.NetID, b.NetID); c != 0 { + return c + } + if c := cmp.Compare(a.Network, b.Network); c != 0 { + return c + } + if c := cmp.Compare(a.Peer, b.Peer); c != 0 { + return c + } + if c := cmp.Compare(a.Metric, b.Metric); c != 0 { + return c + } + return slices.Compare(a.Domains, b.Domains) +} + +func cmpFirewallRule(a, b *proto.FirewallRule) int { + if a == nil || b == nil { + return boolCmp(a == nil, b == nil) + } + if c := bytes.Compare(a.PolicyID, b.PolicyID); c != 0 { + return c + } + if c := cmp.Compare(a.PeerIP, b.PeerIP); c != 0 { //nolint:staticcheck + return c + } + if c := cmp.Compare(int32(a.Direction), int32(b.Direction)); c != 0 { + return c + } + if c := cmp.Compare(int32(a.Action), int32(b.Action)); c != 0 { + return c + } + if c := cmp.Compare(int32(a.Protocol), int32(b.Protocol)); c != 0 { + return c + } + if c := cmp.Compare(a.Port, b.Port); c != 0 { + return c + } + return cmp.Compare(portInfoKey(a.PortInfo), portInfoKey(b.PortInfo)) +} + +func cmpRouteFirewallRule(a, b *proto.RouteFirewallRule) int { + if a == nil || b == nil { + return boolCmp(a == nil, b == nil) + } + if c := bytes.Compare(a.PolicyID, b.PolicyID); c != 0 { + return c + } + if c := cmp.Compare(a.RouteID, b.RouteID); c != 0 { + return c + } + if c := cmp.Compare(a.Destination, b.Destination); c != 0 { + return c + } + if c := cmp.Compare(int32(a.Protocol), int32(b.Protocol)); c != 0 { + return c + } + if c := cmp.Compare(portInfoKey(a.PortInfo), portInfoKey(b.PortInfo)); c != 0 { + return c + } + if c := cmp.Compare(int32(a.Action), int32(b.Action)); c != 0 { + return c + } + if c := slices.Compare(a.Domains, b.Domains); c != 0 { + return c + } + if c := slices.Compare(a.SourceRanges, b.SourceRanges); c != 0 { + return c + } + if c := cmp.Compare(a.CustomProtocol, b.CustomProtocol); c != 0 { + return c + } + return boolCmp(a.IsDynamic, b.IsDynamic) +} + +func cmpForwardingRule(a, b *proto.ForwardingRule) int { + if a == nil || b == nil { + return boolCmp(a == nil, b == nil) + } + if c := cmp.Compare(int32(a.Protocol), int32(b.Protocol)); c != 0 { + return c + } + return bytes.Compare(a.TranslatedAddress, b.TranslatedAddress) +} + +func portInfoKey(pi *proto.PortInfo) string { + if pi == nil { + return "" + } + switch sel := pi.PortSelection.(type) { + case *proto.PortInfo_Port: + return "P" + strconv.FormatUint(uint64(sel.Port), 10) + case *proto.PortInfo_Range_: + if sel.Range == nil { + return "R" + } + return "R" + strconv.FormatUint(uint64(sel.Range.Start), 10) + "-" + strconv.FormatUint(uint64(sel.Range.End), 10) + } + return "" +} + +// describeDivergence names the first differing field so a failure is actionable +// without re-running against the database. +func describeDivergence(legacy, updated *proto.NetworkMap, accountID, peerID string) string { + prefix := "account=" + accountID + " peer=" + peerID + + lens := []struct { + field string + a, b int + }{ + {"RemotePeers", len(legacy.RemotePeers), len(updated.RemotePeers)}, + {"OfflinePeers", len(legacy.OfflinePeers), len(updated.OfflinePeers)}, + {"Routes", len(legacy.Routes), len(updated.Routes)}, + {"FirewallRules", len(legacy.FirewallRules), len(updated.FirewallRules)}, + {"RoutesFirewallRules", len(legacy.RoutesFirewallRules), len(updated.RoutesFirewallRules)}, + {"ForwardingRules", len(legacy.ForwardingRules), len(updated.ForwardingRules)}, + } + for _, l := range lens { + if l.a != l.b { + return prefix + " field=" + l.field + " legacy_len=" + strconv.Itoa(l.a) + " new_len=" + strconv.Itoa(l.b) + } + } + + for i := range legacy.RemotePeers { + if !goproto.Equal(legacy.RemotePeers[i], updated.RemotePeers[i]) { + return prefix + " field=RemotePeers[" + strconv.Itoa(i) + "] legacy=" + protoStr(legacy.RemotePeers[i]) + " new=" + protoStr(updated.RemotePeers[i]) + } + } + for i := range legacy.Routes { + if !goproto.Equal(legacy.Routes[i], updated.Routes[i]) { + return prefix + " field=Routes[" + strconv.Itoa(i) + "] legacy=" + protoStr(legacy.Routes[i]) + " new=" + protoStr(updated.Routes[i]) + } + } + for i := range legacy.FirewallRules { + if !goproto.Equal(legacy.FirewallRules[i], updated.FirewallRules[i]) { + return prefix + " field=FirewallRules[" + strconv.Itoa(i) + "] legacy=" + protoStr(legacy.FirewallRules[i]) + " new=" + protoStr(updated.FirewallRules[i]) + } + } + for i := range legacy.RoutesFirewallRules { + if !goproto.Equal(legacy.RoutesFirewallRules[i], updated.RoutesFirewallRules[i]) { + return prefix + " field=RoutesFirewallRules[" + strconv.Itoa(i) + "] legacy=" + protoStr(legacy.RoutesFirewallRules[i]) + " new=" + protoStr(updated.RoutesFirewallRules[i]) + } + } + if !goproto.Equal(legacy.PeerConfig, updated.PeerConfig) { + return prefix + " field=PeerConfig legacy=" + protoStr(legacy.PeerConfig) + " new=" + protoStr(updated.PeerConfig) + } + if !goproto.Equal(legacy.DNSConfig, updated.DNSConfig) { + return prefix + " field=DNSConfig legacy=" + protoStr(legacy.DNSConfig) + " new=" + protoStr(updated.DNSConfig) + } + if !goproto.Equal(legacy.SshAuth, updated.SshAuth) { + return prefix + " field=SshAuth legacy=" + protoStr(legacy.SshAuth) + " new=" + protoStr(updated.SshAuth) + } + if legacy.Serial != updated.Serial { + return prefix + " field=Serial legacy=" + strconv.FormatUint(legacy.Serial, 10) + " new=" + strconv.FormatUint(updated.Serial, 10) + } + return prefix + " (repeated fields equal element-wise — scalar/oneof mismatch)" +} + +func protoStr(m goproto.Message) string { + if m == nil { + return "" + } + s := prototext.Format(m) + const maxLen = 800 + if len(s) > maxLen { + return s[:maxLen] + "...(truncated)" + } + return s +} diff --git a/management/server/types/legacynmap/firewall_helpers.go b/management/server/types/legacynmap/firewall_helpers.go new file mode 100644 index 000000000..289fc5524 --- /dev/null +++ b/management/server/types/legacynmap/firewall_helpers.go @@ -0,0 +1,157 @@ +//go:build nmapequiv + +package legacynmap + +import ( + "strconv" + "strings" + + v "github.com/hashicorp/go-version" + + "github.com/netbirdio/netbird/version" +) + +const ( + firewallRuleMinPortRangesVer = "0.48.0" + firewallRuleMinNativeSSHVer = "0.60.0" + + nativeSSHPortString = "22022" + nativeSSHPortNumber = 22022 + defaultSSHPortString = "22" + defaultSSHPortNumber = 22 +) + +type supportedFeatures struct { + nativeSSH bool + portRanges bool +} + +type LookupMap map[string]struct{} + +func PolicyRuleImpliesLegacySSH(rule *PolicyRule) bool { + return rule.Protocol == PolicyRuleProtocolALL || (rule.Protocol == PolicyRuleProtocolTCP && (portsIncludesSSH(rule.Ports) || portRangeIncludesSSH(rule.PortRanges))) +} + +func portRangeIncludesSSH(portRanges []RulePortRange) bool { + for _, pr := range portRanges { + if (pr.Start <= defaultSSHPortNumber && pr.End >= defaultSSHPortNumber) || (pr.Start <= nativeSSHPortNumber && pr.End >= nativeSSHPortNumber) { + return true + } + } + return false +} + +func portsIncludesSSH(ports []string) bool { + for _, port := range ports { + if port == defaultSSHPortString || port == nativeSSHPortString { + return true + } + } + return false +} + +// ExpandPortsAndRanges expands Ports and PortRanges of a rule into individual firewall rules. +func ExpandPortsAndRanges(base FirewallRule, rule *PolicyRule, peer *ComponentPeer) []*FirewallRule { + features := peerSupportedFirewallFeatures(peer.AgentVersion) + + var expanded []*FirewallRule + + for _, port := range rule.Ports { + fr := base + fr.Port = port + expanded = append(expanded, &fr) + } + + for _, portRange := range rule.PortRanges { + if len(rule.Ports) > 0 { + break + } + fr := base + + if features.portRanges { + fr.PortRange = portRange + } else { + if portRange.Start != portRange.End { + continue + } + fr.Port = strconv.FormatUint(uint64(portRange.Start), 10) + } + expanded = append(expanded, &fr) + } + + if shouldCheckRulesForNativeSSH(features.nativeSSH, rule, peer) || rule.Protocol == PolicyRuleProtocolNetbirdSSH { + expanded = addNativeSSHRule(base, expanded) + } + + return expanded +} + +func addNativeSSHRule(base FirewallRule, expanded []*FirewallRule) []*FirewallRule { + shouldAdd := false + for _, fr := range expanded { + if isPortInRule(nativeSSHPortString, 22022, fr) { + return expanded + } + if isPortInRule(defaultSSHPortString, 22, fr) { + shouldAdd = true + } + } + if !shouldAdd { + return expanded + } + + fr := base + fr.Port = nativeSSHPortString + return append(expanded, &fr) +} + +func isPortInRule(portString string, portInt uint16, rule *FirewallRule) bool { + return rule.Port == portString || (rule.PortRange.Start <= portInt && portInt <= rule.PortRange.End) +} + +func shouldCheckRulesForNativeSSH(supportsNative bool, rule *PolicyRule, peer *ComponentPeer) bool { + return supportsNative && peer.SSHEnabled && peer.ServerSSHAllowed && rule.Protocol == PolicyRuleProtocolTCP +} + +func peerSupportedFirewallFeatures(peerVer string) supportedFeatures { + if version.IsDevelopmentVersion(peerVer) { + return supportedFeatures{true, true} + } + + var features supportedFeatures + + meetMinVer, err := meetsMinVersion(firewallRuleMinNativeSSHVer, peerVer) + features.nativeSSH = err == nil && meetMinVer + + if features.nativeSSH { + features.portRanges = true + } else { + meetMinVer, err = meetsMinVersion(firewallRuleMinPortRangesVer, peerVer) + features.portRanges = err == nil && meetMinVer + } + + return features +} + +// meetsMinVersion is main's version.MeetsMinVersion, which does not exist at HEAD. +func meetsMinVersion(minVer, peerVer string) (bool, error) { + peerVer = sanitizeVersion(peerVer) + minVer = sanitizeVersion(minVer) + + peerNBVer, err := v.NewVersion(peerVer) + if err != nil { + return false, err + } + + constraints, err := v.NewConstraint(">= " + minVer) + if err != nil { + return false, err + } + + return constraints.Check(peerNBVer), nil +} + +func sanitizeVersion(version string) string { + parts := strings.Split(version, "-") + return parts[0] +} diff --git a/management/server/types/legacynmap/networkmap_components.go b/management/server/types/legacynmap/networkmap_components.go new file mode 100644 index 000000000..64338d727 --- /dev/null +++ b/management/server/types/legacynmap/networkmap_components.go @@ -0,0 +1,1034 @@ +//go:build nmapequiv + +package legacynmap + +import ( + "context" + "maps" + "net/netip" + "slices" + "strconv" + "strings" + "sync" + "time" + + "github.com/netbirdio/netbird/client/ssh/auth" + nbdns "github.com/netbirdio/netbird/dns" + "github.com/netbirdio/netbird/route" + "github.com/netbirdio/netbird/shared/management/domain" +) + +type NetworkMapComponents struct { + PeerID string + + Network *Network + AccountSettings *AccountSettingsInfo + DNSSettings *DNSSettings + CustomZoneDomain string + + Peers map[string]*ComponentPeer + Groups map[string]*ComponentGroup + Policies []*Policy + Routes []*route.Route + NameServerGroups []*nbdns.NameServerGroup + AllDNSRecords []nbdns.SimpleRecord + AccountZones []nbdns.CustomZone + ResourcePoliciesMap map[string][]*Policy + RoutersMap map[string]map[string]*ComponentRouter + NetworkResources []*ComponentResource + + GroupIDToUserIDs map[string][]string + AllowedUserIDs map[string]struct{} + PostureFailedPeers map[string]map[string]struct{} + + RouterPeers map[string]*ComponentPeer + + // NetworkXIDToPublicID maps Network.ID (xid) → PublicID. + // Consumed by the envelope encoder to + // translate RoutersMap keys and NetworkResource.NetworkID references + // to compact uint32 ids. Legacy Calculate() doesn't consult it. + NetworkXIDToPublicID map[string]string + + // PostureCheckXIDToPublicID maps posture.Checks.ID (xid) → PublicID. + // Same role as NetworkXIDToPublicID, used for PostureFailedPeers keys and + // policy SourcePostureChecks references. + PostureCheckXIDToPublicID map[string]string + routesByPeerOnce sync.Once + routesByPeerIdx map[string][]routeIndexEntry + + // true when returning an empty-like map (returned instead of nil) + empty bool + + // ForceRoutingPeerDNSResolution forces the peer to run/use routing-peer DNS + // resolution regardless of the account-global setting, for reverse-proxy + // domain targets. + ForceRoutingPeerDNSResolution bool +} + +type routeIndexEntry struct { + route *route.Route + viaGroup bool +} + +type AccountSettingsInfo struct { + PeerLoginExpirationEnabled bool + PeerLoginExpiration time.Duration + PeerInactivityExpirationEnabled bool + PeerInactivityExpiration time.Duration +} + +func EmptyNetworkMapComponents(nm *NetworkMapComponents) *NetworkMapComponents { + nm.empty = true + return nm +} + +func (c *NetworkMapComponents) GetPeerInfo(peerID string) *ComponentPeer { + return c.Peers[peerID] +} + +func (c *NetworkMapComponents) GetRouterPeerInfo(peerID string) *ComponentPeer { + return c.RouterPeers[peerID] +} + +func (c *NetworkMapComponents) GetGroupInfo(groupID string) *ComponentGroup { + return c.Groups[groupID] +} + +func (c *NetworkMapComponents) IsPeerInGroup(peerID, groupID string) bool { + group := c.GetGroupInfo(groupID) + if group == nil { + return false + } + + return slices.Contains(group.Peers, peerID) +} + +func (c *NetworkMapComponents) GetPeerGroups(peerID string) map[string]struct{} { + groups := make(map[string]struct{}) + for groupID, group := range c.Groups { + if slices.Contains(group.Peers, peerID) { + groups[groupID] = struct{}{} + } + } + return groups +} + +func (c *NetworkMapComponents) ValidatePostureChecksOnPeer(peerID string, postureCheckIDs []string) bool { + _, exists := c.Peers[peerID] + if !exists { + return false + } + if len(postureCheckIDs) == 0 { + return true + } + for _, checkID := range postureCheckIDs { + if failedPeers, exists := c.PostureFailedPeers[checkID]; exists { + if _, failed := failedPeers[peerID]; failed { + return false + } + } + } + return true +} + +func CalculateNetworkMapFromComponents(ctx context.Context, components *NetworkMapComponents) *NetworkMap { + return components.Calculate(ctx) +} + +func (c *NetworkMapComponents) Calculate(ctx context.Context) *NetworkMap { + targetPeerID := c.PeerID + + peerGroups := c.GetPeerGroups(targetPeerID) + + aclPeers, firewallRules, authorizedUsers, sshEnabled := c.getPeerConnectionResources(targetPeerID) + + peersToConnect, expiredPeers := c.filterPeersByLoginExpiration(aclPeers) + + includeIPv6 := false + if p := c.Peers[targetPeerID]; p != nil { + includeIPv6 = p.SupportsIPv6 && p.IPv6.IsValid() + } + routesUpdate := filterAndExpandRoutes(c.getRoutesToSync(targetPeerID, peersToConnect, peerGroups), includeIPv6) + routesFirewallRules := c.getPeerRoutesFirewallRules(ctx, targetPeerID, includeIPv6) + + isRouter, networkResourcesRoutes, sourcePeers := c.getNetworkResourcesRoutesToSync(targetPeerID) + var networkResourcesFirewallRules []*RouteFirewallRule + if isRouter { + networkResourcesFirewallRules = c.getPeerNetworkResourceFirewallRules(ctx, targetPeerID, networkResourcesRoutes, includeIPv6) + } + + peersToConnectIncludingRouters := c.addNetworksRoutingPeers( + networkResourcesRoutes, + targetPeerID, + peersToConnect, + expiredPeers, + isRouter, + sourcePeers, + ) + + dnsManagementStatus := c.getPeerDNSManagementStatusFromGroups(peerGroups) + dnsUpdate := nbdns.Config{ + ServiceEnable: dnsManagementStatus, + } + + if dnsManagementStatus { + var customZones []nbdns.CustomZone + + if c.CustomZoneDomain != "" && len(c.AllDNSRecords) > 0 { + customZones = append(customZones, nbdns.CustomZone{ + Domain: c.CustomZoneDomain, + Records: c.AllDNSRecords, + }) + } + + customZones = append(customZones, c.AccountZones...) + + dnsUpdate.CustomZones = customZones + dnsUpdate.NameServerGroups = c.getPeerNSGroupsFromGroups(targetPeerID, peerGroups) + } + + return &NetworkMap{ + Peers: peersToConnectIncludingRouters, + Network: c.Network.Copy(), + Routes: append(filterAndExpandRoutes(networkResourcesRoutes, includeIPv6), routesUpdate...), + DNSConfig: dnsUpdate, + OfflinePeers: expiredPeers, + FirewallRules: firewallRules, + RoutesFirewallRules: append(networkResourcesFirewallRules, routesFirewallRules...), + AuthorizedUsers: authorizedUsers, + EnableSSH: sshEnabled, + + ForceRoutingPeerDNSResolution: c.ForceRoutingPeerDNSResolution, + } +} + +func (c *NetworkMapComponents) IsEmpty() bool { + return c.empty +} + +func (c *NetworkMapComponents) getPeerConnectionResources(targetPeerID string) ([]*ComponentPeer, []*FirewallRule, map[string]map[string]struct{}, bool) { + targetPeer := c.GetPeerInfo(targetPeerID) + if targetPeer == nil { + return nil, nil, nil, false + } + + generateResources, getAccumulatedResources := c.connResourcesGenerator(targetPeer) + authorizedUsers := make(map[string]map[string]struct{}) + sshEnabled := false + + for _, policy := range c.Policies { + if !policy.Enabled { + continue + } + + for _, rule := range policy.Rules { + if !rule.Enabled { + continue + } + + var sourcePeers, destinationPeers []*ComponentPeer + var peerInSources, peerInDestinations bool + + if rule.SourceResource.Type == ResourceTypePeer && rule.SourceResource.ID != "" { + sourcePeers, peerInSources = c.getPeerFromResource(rule.SourceResource, targetPeerID) + } else { + sourcePeers, peerInSources = c.getAllPeersFromGroups(rule.Sources, targetPeerID, policy.SourcePostureChecks) + } + + if rule.DestinationResource.Type == ResourceTypePeer && rule.DestinationResource.ID != "" { + destinationPeers, peerInDestinations = c.getPeerFromResource(rule.DestinationResource, targetPeerID) + } else { + destinationPeers, peerInDestinations = c.getAllPeersFromGroups(rule.Destinations, targetPeerID, nil) + } + + if rule.Bidirectional { + if peerInSources { + generateResources(rule, destinationPeers, FirewallRuleDirectionIN) + } + if peerInDestinations { + generateResources(rule, sourcePeers, FirewallRuleDirectionOUT) + } + } + + if peerInSources { + generateResources(rule, destinationPeers, FirewallRuleDirectionOUT) + } + + if peerInDestinations { + generateResources(rule, sourcePeers, FirewallRuleDirectionIN) + } + + if peerInDestinations && rule.Protocol == PolicyRuleProtocolNetbirdSSH { + sshEnabled = true + switch { + case len(rule.AuthorizedGroups) > 0: + for groupID, localUsers := range rule.AuthorizedGroups { + userIDs, ok := c.GroupIDToUserIDs[groupID] + if !ok { + continue + } + + if len(localUsers) == 0 { + localUsers = []string{auth.Wildcard} + } + + for _, localUser := range localUsers { + if authorizedUsers[localUser] == nil { + authorizedUsers[localUser] = make(map[string]struct{}) + } + for _, userID := range userIDs { + authorizedUsers[localUser][userID] = struct{}{} + } + } + } + case rule.AuthorizedUser != "": + if authorizedUsers[auth.Wildcard] == nil { + authorizedUsers[auth.Wildcard] = make(map[string]struct{}) + } + authorizedUsers[auth.Wildcard][rule.AuthorizedUser] = struct{}{} + default: + authorizedUsers[auth.Wildcard] = c.getAllowedUserIDs() + } + } else if peerInDestinations && PolicyRuleImpliesLegacySSH(rule) && targetPeer.SSHEnabled { + sshEnabled = true + authorizedUsers[auth.Wildcard] = c.getAllowedUserIDs() + } + } + } + + peers, fwRules := getAccumulatedResources() + return peers, fwRules, authorizedUsers, sshEnabled +} + +func (c *NetworkMapComponents) getAllowedUserIDs() map[string]struct{} { + if c.AllowedUserIDs != nil { + result := make(map[string]struct{}, len(c.AllowedUserIDs)) + maps.Copy(result, c.AllowedUserIDs) + return result + } + return make(map[string]struct{}) +} + +func (c *NetworkMapComponents) connResourcesGenerator(targetPeer *ComponentPeer) (func(*PolicyRule, []*ComponentPeer, int), func() ([]*ComponentPeer, []*FirewallRule)) { + rulesExists := make(map[string]struct{}) + peersExists := make(map[string]struct{}) + rules := make([]*FirewallRule, 0) + peers := make([]*ComponentPeer, 0) + + return func(rule *PolicyRule, groupPeers []*ComponentPeer, direction int) { + protocol := rule.Protocol + if protocol == PolicyRuleProtocolNetbirdSSH { + protocol = PolicyRuleProtocolTCP + } + + protocolStr := string(protocol) + actionStr := string(rule.Action) + dirStr := strconv.Itoa(direction) + portsJoined := strings.Join(rule.Ports, ",") + + for _, peer := range groupPeers { + if peer == nil { + continue + } + + if _, ok := peersExists[peer.ID]; !ok { + peers = append(peers, peer) + peersExists[peer.ID] = struct{}{} + } + + peerIP := peer.IP.String() + + fr := FirewallRule{ + PolicyID: rule.ID, + PeerIP: peerIP, + Direction: direction, + Action: actionStr, + Protocol: protocolStr, + } + + ruleID := rule.ID + peerIP + dirStr + + protocolStr + actionStr + portsJoined + if _, ok := rulesExists[ruleID]; ok { + continue + } + rulesExists[ruleID] = struct{}{} + + if len(rule.Ports) == 0 && len(rule.PortRanges) == 0 { + rules = append(rules, &fr) + } else { + rules = append(rules, ExpandPortsAndRanges(fr, rule, targetPeer)...) + } + + rules = AppendIPv6FirewallRule(rules, rulesExists, peer, targetPeer, rule, FirewallRuleContext{ + Direction: direction, + DirStr: dirStr, + ProtocolStr: protocolStr, + ActionStr: actionStr, + PortsJoined: portsJoined, + }) + } + }, func() ([]*ComponentPeer, []*FirewallRule) { + return peers, rules + } +} + +func (c *NetworkMapComponents) getAllPeersFromGroups(groups []string, peerID string, sourcePostureChecksIDs []string) ([]*ComponentPeer, bool) { + peerInGroups := false + uniquePeerIDs := c.getUniquePeerIDsFromGroupsIDs(groups) + filteredPeers := make([]*ComponentPeer, 0, len(uniquePeerIDs)) + + for _, p := range uniquePeerIDs { + peerInfo := c.GetPeerInfo(p) + if peerInfo == nil { + continue + } + + if _, ok := c.Peers[p]; !ok { + continue + } + + if !c.ValidatePostureChecksOnPeer(p, sourcePostureChecksIDs) { + continue + } + + if p == peerID { + peerInGroups = true + continue + } + + filteredPeers = append(filteredPeers, peerInfo) + } + + return filteredPeers, peerInGroups +} + +func (c *NetworkMapComponents) getUniquePeerIDsFromGroupsIDs(groups []string) []string { + peerIDs := make(map[string]struct{}, len(groups)) + for _, groupID := range groups { + group := c.GetGroupInfo(groupID) + if group == nil { + continue + } + + if group.IsGroupAll() || len(groups) == 1 { + return group.Peers + } + + for _, peerID := range group.Peers { + peerIDs[peerID] = struct{}{} + } + } + + ids := make([]string, 0, len(peerIDs)) + for peerID := range peerIDs { + ids = append(ids, peerID) + } + + return ids +} + +func (c *NetworkMapComponents) getPeerFromResource(resource Resource, peerID string) ([]*ComponentPeer, bool) { + if resource.ID == peerID { + return []*ComponentPeer{}, true + } + + peerInfo := c.GetPeerInfo(resource.ID) + if peerInfo == nil { + return []*ComponentPeer{}, false + } + + return []*ComponentPeer{peerInfo}, false +} + +func (c *NetworkMapComponents) filterPeersByLoginExpiration(aclPeers []*ComponentPeer) ([]*ComponentPeer, []*ComponentPeer) { + peersToConnect := make([]*ComponentPeer, 0, len(aclPeers)) + var expiredPeers []*ComponentPeer + + for _, p := range aclPeers { + expired, _ := p.LoginExpired(c.AccountSettings.PeerLoginExpiration) + if c.AccountSettings.PeerLoginExpirationEnabled && expired { + expiredPeers = append(expiredPeers, p) + continue + } + peersToConnect = append(peersToConnect, p) + } + + return peersToConnect, expiredPeers +} + +func (c *NetworkMapComponents) getPeerDNSManagementStatusFromGroups(peerGroups map[string]struct{}) bool { + for _, groupID := range c.DNSSettings.DisabledManagementGroups { + if _, found := peerGroups[groupID]; found { + return false + } + } + return true +} + +func (c *NetworkMapComponents) getPeerNSGroupsFromGroups(peerID string, groupList map[string]struct{}) []*nbdns.NameServerGroup { + var peerNSGroups []*nbdns.NameServerGroup + + targetPeerInfo := c.GetPeerInfo(peerID) + if targetPeerInfo == nil { + return peerNSGroups + } + + peerIPStr := targetPeerInfo.IP.String() + + for _, nsGroup := range c.NameServerGroups { + if !nsGroup.Enabled { + continue + } + for _, gID := range nsGroup.Groups { + if _, found := groupList[gID]; found { + if !c.peerIsNameserver(peerIPStr, nsGroup) { + peerNSGroups = append(peerNSGroups, nsGroup.Copy()) + } + break + } + } + } + + return peerNSGroups +} + +func (c *NetworkMapComponents) peerIsNameserver(peerIPStr string, nsGroup *nbdns.NameServerGroup) bool { + for _, ns := range nsGroup.NameServers { + if peerIPStr == ns.IP.String() { + return true + } + } + return false +} + +// filterAndExpandRoutes drops v6 routes for non-capable peers and duplicates +// the default v4 route (0.0.0.0/0) as ::/0 for v6-capable peers. +// TODO: the "-v6" suffix on IDs could collide with user-supplied route IDs. +func filterAndExpandRoutes(routes []*route.Route, includeIPv6 bool) []*route.Route { + filtered := make([]*route.Route, 0, len(routes)) + for _, r := range routes { + if !includeIPv6 && r.Network.Addr().Is6() { + continue + } + filtered = append(filtered, r) + + if includeIPv6 && r.Network.Bits() == 0 && r.Network.Addr().Is4() { + v6 := r.Copy() + v6.ID = r.ID + "-v6-default" + v6.NetID = r.NetID + "-v6" + v6.Network = netip.MustParsePrefix("::/0") + v6.NetworkType = route.IPv6Network + filtered = append(filtered, v6) + } + } + return filtered +} + +func (c *NetworkMapComponents) getRoutesToSync(peerID string, aclPeers []*ComponentPeer, peerGroups LookupMap) []*route.Route { + routes, peerDisabledRoutes := c.getRoutingPeerRoutes(peerID) + peerRoutesMembership := make(LookupMap) + for _, r := range append(routes, peerDisabledRoutes...) { + peerRoutesMembership[string(r.GetHAUniqueID())] = struct{}{} + } + + for _, peer := range aclPeers { + activeRoutes, _ := c.getRoutingPeerRoutes(peer.ID) + groupFilteredRoutes := c.filterRoutesByGroups(activeRoutes, peerGroups) + filteredRoutes := c.filterRoutesFromPeersOfSameHAGroup(groupFilteredRoutes, peerRoutesMembership) + routes = append(routes, filteredRoutes...) + } + + return routes +} + +func (c *NetworkMapComponents) getRoutingPeerRoutes(peerID string) (enabledRoutes []*route.Route, disabledRoutes []*route.Route) { + peerInfo := c.GetPeerInfo(peerID) + if peerInfo == nil { + peerInfo = c.GetRouterPeerInfo(peerID) + } + if peerInfo == nil { + return enabledRoutes, disabledRoutes + } + + seenRoute := make(map[route.ID]struct{}) + + takeRoute := func(r *route.Route) { + if _, ok := seenRoute[r.ID]; ok { + return + } + seenRoute[r.ID] = struct{}{} + + r.Peer = peerInfo.Key + + if r.Enabled { + enabledRoutes = append(enabledRoutes, r) + return + } + disabledRoutes = append(disabledRoutes, r) + } + + for _, entry := range c.routesByPeer()[peerID] { + if entry.viaGroup { + newPeerRoute := entry.route.Copy() + newPeerRoute.PeerGroups = nil + newPeerRoute.ID = route.ID(string(entry.route.ID) + ":" + peerID) + takeRoute(newPeerRoute) + continue + } + takeRoute(entry.route.Copy()) + } + + return enabledRoutes, disabledRoutes +} + +func (c *NetworkMapComponents) routesByPeer() map[string][]routeIndexEntry { + c.routesByPeerOnce.Do(func() { + idx := make(map[string][]routeIndexEntry) + for _, r := range c.Routes { + for _, groupID := range r.PeerGroups { + group := c.GetGroupInfo(groupID) + if group == nil { + continue + } + for _, id := range group.Peers { + idx[id] = append(idx[id], routeIndexEntry{route: r, viaGroup: true}) + } + } + if r.Peer != "" { + idx[r.Peer] = append(idx[r.Peer], routeIndexEntry{route: r}) + } + } + c.routesByPeerIdx = idx + }) + + return c.routesByPeerIdx +} + +func (c *NetworkMapComponents) filterRoutesByGroups(routes []*route.Route, groupListMap LookupMap) []*route.Route { + var filteredRoutes []*route.Route + for _, r := range routes { + for _, groupID := range r.Groups { + _, found := groupListMap[groupID] + if found { + filteredRoutes = append(filteredRoutes, r) + break + } + } + } + return filteredRoutes +} + +func (c *NetworkMapComponents) filterRoutesFromPeersOfSameHAGroup(routes []*route.Route, peerMemberships LookupMap) []*route.Route { + var filteredRoutes []*route.Route + for _, r := range routes { + _, found := peerMemberships[string(r.GetHAUniqueID())] + if !found { + filteredRoutes = append(filteredRoutes, r) + } + } + return filteredRoutes +} + +func (c *NetworkMapComponents) getPeerRoutesFirewallRules(ctx context.Context, peerID string, includeIPv6 bool) []*RouteFirewallRule { + routesFirewallRules := make([]*RouteFirewallRule, 0) + + enabledRoutes, _ := c.getRoutingPeerRoutes(peerID) + for _, r := range enabledRoutes { + if len(r.AccessControlGroups) == 0 { + defaultPermit := c.getDefaultPermit(r, includeIPv6) + routesFirewallRules = append(routesFirewallRules, defaultPermit...) + continue + } + + distributionPeers := c.getDistributionGroupsPeers(r) + + for _, accessGroup := range r.AccessControlGroups { + policies := c.getAllRoutePoliciesFromGroups([]string{accessGroup}) + rules := c.getRouteFirewallRules(ctx, peerID, policies, r, distributionPeers, includeIPv6) + routesFirewallRules = append(routesFirewallRules, rules...) + } + } + + return routesFirewallRules +} + +func (c *NetworkMapComponents) getDefaultPermit(r *route.Route, includeIPv6 bool) []*RouteFirewallRule { + if r.Network.Addr().Is6() && !includeIPv6 { + return nil + } + + sources := []string{"0.0.0.0/0"} + if r.Network.Addr().Is6() { + sources = []string{"::/0"} + } + + rule := RouteFirewallRule{ + SourceRanges: sources, + Action: string(PolicyTrafficActionAccept), + Destination: r.Network.String(), + Protocol: string(PolicyRuleProtocolALL), + Domains: r.Domains, + IsDynamic: r.IsDynamic(), + RouteID: r.ID, + } + + rules := []*RouteFirewallRule{&rule} + + isDefaultV4 := r.Network.Addr().Is4() && r.Network.Bits() == 0 + if includeIPv6 && (r.IsDynamic() || isDefaultV4) { + ruleV6 := rule + ruleV6.SourceRanges = []string{"::/0"} + if isDefaultV4 { + ruleV6.Destination = "::/0" + ruleV6.RouteID = r.ID + "-v6-default" + } + rules = append(rules, &ruleV6) + } + + return rules +} + +func (c *NetworkMapComponents) getDistributionGroupsPeers(r *route.Route) map[string]struct{} { + distPeers := make(map[string]struct{}) + for _, id := range r.Groups { + group := c.GetGroupInfo(id) + if group == nil { + continue + } + + for _, pID := range group.Peers { + distPeers[pID] = struct{}{} + } + } + return distPeers +} + +func (c *NetworkMapComponents) getAllRoutePoliciesFromGroups(accessControlGroups []string) []*Policy { + routePolicies := make([]*Policy, 0) + for _, groupID := range accessControlGroups { + for _, policy := range c.Policies { + for _, rule := range policy.Rules { + if slices.Contains(rule.Destinations, groupID) { + routePolicies = append(routePolicies, policy) + } + } + } + } + + return routePolicies +} + +func (c *NetworkMapComponents) getRouteFirewallRules(ctx context.Context, peerID string, policies []*Policy, route *route.Route, distributionPeers map[string]struct{}, includeIPv6 bool) []*RouteFirewallRule { + var fwRules []*RouteFirewallRule + for _, policy := range policies { + if !policy.Enabled { + continue + } + + for _, rule := range policy.Rules { + if !rule.Enabled { + continue + } + + rulePeers := c.getRulePeers(rule, policy.SourcePostureChecks, peerID, distributionPeers) + rules := GenerateRouteFirewallRules(ctx, route, rule, rulePeers, FirewallRuleDirectionIN, includeIPv6) + fwRules = append(fwRules, rules...) + } + } + return fwRules +} + +func (c *NetworkMapComponents) getRulePeers(rule *PolicyRule, postureChecks []string, peerID string, distributionPeers map[string]struct{}) []*ComponentPeer { + distPeersWithPolicy := make(map[string]struct{}) + for _, id := range rule.Sources { + group := c.GetGroupInfo(id) + if group == nil { + continue + } + + for _, pID := range group.Peers { + if pID == peerID { + continue + } + _, distPeer := distributionPeers[pID] + _, valid := c.Peers[pID] + if distPeer && valid && c.ValidatePostureChecksOnPeer(pID, postureChecks) { + distPeersWithPolicy[pID] = struct{}{} + } + } + } + if rule.SourceResource.Type == ResourceTypePeer && rule.SourceResource.ID != "" { + _, distPeer := distributionPeers[rule.SourceResource.ID] + _, valid := c.Peers[rule.SourceResource.ID] + if distPeer && valid && c.ValidatePostureChecksOnPeer(rule.SourceResource.ID, postureChecks) { + distPeersWithPolicy[rule.SourceResource.ID] = struct{}{} + } + } + + distributionGroupPeers := make([]*ComponentPeer, 0, len(distPeersWithPolicy)) + for pID := range distPeersWithPolicy { + peerInfo := c.GetPeerInfo(pID) + if peerInfo == nil { + continue + } + distributionGroupPeers = append(distributionGroupPeers, peerInfo) + } + return distributionGroupPeers +} + +func (c *NetworkMapComponents) getNetworkResourcesRoutesToSync(peerID string) (bool, []*route.Route, map[string]struct{}) { + var isRoutingPeer bool + var routes []*route.Route + allSourcePeers := make(map[string]struct{}) + + for _, resource := range c.NetworkResources { + if !resource.Enabled { + continue + } + + var addSourcePeers bool + + networkRoutingPeers, exists := c.RoutersMap[resource.NetworkID] + if exists { + if router, ok := networkRoutingPeers[peerID]; ok { + isRoutingPeer, addSourcePeers = true, true + routes = append(routes, c.getNetworkResourcesRoutes(resource, peerID, router)...) + } + } + + newRoutes := c.processResourcePolicies(peerID, resource, networkRoutingPeers, addSourcePeers, allSourcePeers) + routes = append(routes, newRoutes...) + } + + return isRoutingPeer, routes, allSourcePeers +} + +func (c *NetworkMapComponents) processResourcePolicies( + peerID string, + resource *ComponentResource, + networkRoutingPeers map[string]*ComponentRouter, + addSourcePeers bool, + allSourcePeers map[string]struct{}, +) []*route.Route { + var routes []*route.Route + + for _, policy := range c.ResourcePoliciesMap[resource.ID] { + peers := c.getResourcePolicyPeers(policy) + if addSourcePeers { + for _, pID := range c.getPostureValidPeers(peers, policy.SourcePostureChecks) { + allSourcePeers[pID] = struct{}{} + } + continue + } + + if slices.Contains(peers, peerID) && c.ValidatePostureChecksOnPeer(peerID, policy.SourcePostureChecks) { + for peerId, router := range networkRoutingPeers { + routes = append(routes, c.getNetworkResourcesRoutes(resource, peerId, router)...) + } + break + } + } + + return routes +} + +func (c *NetworkMapComponents) getResourcePolicyPeers(policy *Policy) []string { + if policy.Rules[0].SourceResource.Type == ResourceTypePeer && policy.Rules[0].SourceResource.ID != "" { + return []string{policy.Rules[0].SourceResource.ID} + } + return c.getUniquePeerIDsFromGroupsIDs(policy.SourceGroups()) +} + +func (c *NetworkMapComponents) getNetworkResourcesRoutes(resource *ComponentResource, peerID string, router *ComponentRouter) []*route.Route { + resourceAppliedPolicies := c.ResourcePoliciesMap[resource.ID] + + var routes []*route.Route + if len(resourceAppliedPolicies) > 0 { + peerInfo := c.GetPeerInfo(peerID) + if peerInfo != nil { + routes = append(routes, c.networkResourceToRoute(resource, peerInfo, router)) + } + } + + return routes +} + +func (c *NetworkMapComponents) networkResourceToRoute(resource *ComponentResource, peer *ComponentPeer, router *ComponentRouter) *route.Route { + r := &route.Route{ + ID: route.ID(resource.ID + ":" + peer.ID), + AccountID: resource.AccountID, + Peer: peer.Key, + PeerID: peer.ID, + Metric: router.Metric, + Masquerade: router.Masquerade, + Enabled: resource.Enabled, + KeepRoute: true, + NetID: route.NetID(resource.Name), + Description: resource.Description, + } + + if resource.Type == ComponentResourceHost || resource.Type == ComponentResourceSubnet { + r.Network = resource.Prefix + + r.NetworkType = route.IPv4Network + if resource.Prefix.Addr().Is6() { + r.NetworkType = route.IPv6Network + } + } + + if resource.Type == ComponentResourceDomain { + domainList, err := domain.FromStringList([]string{resource.Domain}) + if err == nil { + r.Domains = domainList + r.NetworkType = route.DomainNetwork + r.Network = netip.PrefixFrom(netip.AddrFrom4([4]byte{192, 0, 2, 0}), 32) + } + } + + return r +} + +func (c *NetworkMapComponents) getPostureValidPeers(inputPeers []string, postureChecksIDs []string) []string { + var dest []string + for _, peerID := range inputPeers { + if c.ValidatePostureChecksOnPeer(peerID, postureChecksIDs) { + dest = append(dest, peerID) + } + } + return dest +} + +func (c *NetworkMapComponents) getPeerNetworkResourceFirewallRules(ctx context.Context, peerID string, routes []*route.Route, includeIPv6 bool) []*RouteFirewallRule { + routesFirewallRules := make([]*RouteFirewallRule, 0) + + peerInfo := c.GetPeerInfo(peerID) + if peerInfo == nil { + return routesFirewallRules + } + + for _, r := range routes { + if r.Peer != peerInfo.Key { + continue + } + + resourceID := string(r.GetResourceID()) + resourcePolicies := c.ResourcePoliciesMap[resourceID] + distributionPeers := c.getPoliciesSourcePeers(resourcePolicies) + + rules := c.getRouteFirewallRules(ctx, peerID, resourcePolicies, r, distributionPeers, includeIPv6) + for _, rule := range rules { + if len(rule.SourceRanges) > 0 { + routesFirewallRules = append(routesFirewallRules, rule) + } + } + } + + return routesFirewallRules +} + +func (c *NetworkMapComponents) getPoliciesSourcePeers(policies []*Policy) map[string]struct{} { + sourcePeers := make(map[string]struct{}) + + for _, policy := range policies { + for _, rule := range policy.Rules { + for _, sourceGroup := range rule.Sources { + group := c.GetGroupInfo(sourceGroup) + if group == nil { + continue + } + + for _, peer := range group.Peers { + sourcePeers[peer] = struct{}{} + } + } + + if rule.SourceResource.Type == ResourceTypePeer && rule.SourceResource.ID != "" { + sourcePeers[rule.SourceResource.ID] = struct{}{} + } + } + } + + return sourcePeers +} + +func (c *NetworkMapComponents) addNetworksRoutingPeers( + networkResourcesRoutes []*route.Route, + peerID string, + peersToConnect []*ComponentPeer, + expiredPeers []*ComponentPeer, + isRouter bool, + sourcePeers map[string]struct{}, +) []*ComponentPeer { + + networkRoutesPeers := make(map[string]struct{}, len(networkResourcesRoutes)) + for _, r := range networkResourcesRoutes { + networkRoutesPeers[r.PeerID] = struct{}{} + } + + delete(sourcePeers, peerID) + delete(networkRoutesPeers, peerID) + + for _, existingPeer := range peersToConnect { + delete(sourcePeers, existingPeer.ID) + delete(networkRoutesPeers, existingPeer.ID) + } + for _, expPeer := range expiredPeers { + delete(sourcePeers, expPeer.ID) + delete(networkRoutesPeers, expPeer.ID) + } + + missingPeers := make(map[string]struct{}, len(sourcePeers)+len(networkRoutesPeers)) + if isRouter { + for p := range sourcePeers { + missingPeers[p] = struct{}{} + } + } + for p := range networkRoutesPeers { + missingPeers[p] = struct{}{} + } + + for p := range missingPeers { + peerInfo := c.GetPeerInfo(p) + if peerInfo == nil { + peerInfo = c.GetRouterPeerInfo(p) + } + if peerInfo != nil { + peersToConnect = append(peersToConnect, peerInfo) + } + } + + return peersToConnect +} + +type FirewallRuleContext struct { + Direction int + DirStr string + ProtocolStr string + ActionStr string + PortsJoined string +} + +func AppendIPv6FirewallRule(rules []*FirewallRule, rulesExists map[string]struct{}, peer, targetPeer *ComponentPeer, rule *PolicyRule, rc FirewallRuleContext) []*FirewallRule { + if !peer.IPv6.IsValid() || !targetPeer.SupportsIPv6 || !targetPeer.IPv6.IsValid() { + return rules + } + + v6IP := peer.IPv6.String() + v6RuleID := rule.ID + v6IP + rc.DirStr + rc.ProtocolStr + rc.ActionStr + rc.PortsJoined + if _, ok := rulesExists[v6RuleID]; ok { + return rules + } + rulesExists[v6RuleID] = struct{}{} + + v6fr := FirewallRule{ + PolicyID: rule.ID, + PeerIP: v6IP, + Direction: rc.Direction, + Action: rc.ActionStr, + Protocol: rc.ProtocolStr, + } + if len(rule.Ports) == 0 && len(rule.PortRanges) == 0 { + return append(rules, &v6fr) + } + return append(rules, ExpandPortsAndRanges(v6fr, rule, targetPeer)...) +} diff --git a/management/server/types/legacynmap/proto_legacy.go b/management/server/types/legacynmap/proto_legacy.go new file mode 100644 index 000000000..6bd673bfd --- /dev/null +++ b/management/server/types/legacynmap/proto_legacy.go @@ -0,0 +1,208 @@ +//go:build nmapequiv + +package legacynmap + +import ( + "context" + "fmt" + "net/netip" + "net/url" + "strings" + + "github.com/netbirdio/netbird/client/ssh/auth" + nbconfig "github.com/netbirdio/netbird/management/internals/server/config" + nbpeer "github.com/netbirdio/netbird/management/server/peer" + types "github.com/netbirdio/netbird/management/server/types" + nbroute "github.com/netbirdio/netbird/route" + "github.com/netbirdio/netbird/shared/management/networkmap" + "github.com/netbirdio/netbird/shared/management/proto" + "github.com/netbirdio/netbird/shared/netiputil" +) + +func ToProtocolRoutes(routes []*nbroute.Route) []*proto.Route { + protoRoutes := make([]*proto.Route, 0, len(routes)) + for _, r := range routes { + protoRoutes = append(protoRoutes, ToProtocolRoute(r)) + } + return protoRoutes +} + +func ToProtocolRoute(route *nbroute.Route) *proto.Route { + return &proto.Route{ + ID: string(route.ID), + NetID: string(route.NetID), + Network: route.Network.String(), + Domains: route.Domains.ToPunycodeList(), + NetworkType: int64(route.NetworkType), + Peer: route.Peer, + Metric: int64(route.Metric), + Masquerade: route.Masquerade, + KeepRoute: route.KeepRoute, + SkipAutoApply: route.SkipAutoApply, + } +} + +func AppendRemotePeerConfig(dst []*proto.RemotePeerConfig, peers []*ComponentPeer, dnsName string, includeIPv6 bool) []*proto.RemotePeerConfig { + for _, rPeer := range peers { + allowedIPs := []string{rPeer.IP.String() + "/32"} + if includeIPv6 && rPeer.IPv6.IsValid() { + allowedIPs = append(allowedIPs, rPeer.IPv6.String()+"/128") + } + dst = append(dst, &proto.RemotePeerConfig{ + WgPubKey: rPeer.Key, + AllowedIps: allowedIPs, + SshConfig: &proto.SSHConfig{SshPubKey: []byte(rPeer.SSHKey)}, + Fqdn: rPeer.FQDN(dnsName), + AgentVersion: rPeer.AgentVersion, + }) + } + return dst +} + +func buildJWTConfig(config *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow) *proto.JWTConfig { + if config == nil || config.AuthAudience == "" { + return nil + } + + issuer := strings.TrimSpace(config.AuthIssuer) + if issuer == "" && deviceFlowConfig != nil { + if d := deriveIssuerFromTokenEndpoint(deviceFlowConfig.ProviderConfig.TokenEndpoint); d != "" { + issuer = d + } + } + if issuer == "" { + return nil + } + + keysLocation := strings.TrimSpace(config.AuthKeysLocation) + if keysLocation == "" { + keysLocation = strings.TrimSuffix(issuer, "/") + "/.well-known/jwks.json" + } + + audience := config.AuthAudience + if config.CLIAuthAudience != "" { + audience = config.CLIAuthAudience + } + + audiences := []string{config.AuthAudience} + if config.CLIAuthAudience != "" && config.CLIAuthAudience != config.AuthAudience { + audiences = append(audiences, config.CLIAuthAudience) + } + + return &proto.JWTConfig{ + Issuer: issuer, + Audience: audience, + Audiences: audiences, + KeysLocation: keysLocation, + } +} + +func toPeerConfig(peer *nbpeer.Peer, network *Network, dnsName string, settings *types.Settings, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, enableSSH bool, forceRoutingPeerDNS bool) *proto.PeerConfig { + netmask, _ := network.Net.Mask.Size() + fqdn := peer.FQDN(dnsName) + + sshConfig := &proto.SSHConfig{ + SshEnabled: peer.SSHEnabled || enableSSH, + } + + if sshConfig.SshEnabled { + sshConfig.JwtConfig = buildJWTConfig(httpConfig, deviceFlowConfig) + } + + peerConfig := &proto.PeerConfig{ + Address: fmt.Sprintf("%s/%d", peer.IP.String(), netmask), + SshConfig: sshConfig, + Fqdn: fqdn, + RoutingPeerDnsResolutionEnabled: settings.RoutingPeerDNSResolutionEnabled || peer.ProxyMeta.Embedded || forceRoutingPeerDNS, + LazyConnectionEnabled: settings.LazyConnectionEnabled, + AutoUpdate: &proto.AutoUpdateSettings{ + Version: settings.AutoUpdateVersion, + AlwaysUpdate: settings.AutoUpdateAlways, + }, + } + + if peer.SupportsIPv6() && peer.IPv6.IsValid() && network.NetV6.IP != nil { + ones, _ := network.NetV6.Mask.Size() + v6Prefix := netip.PrefixFrom(peer.IPv6.Unmap(), ones) + if b, err := netiputil.EncodePrefix(v6Prefix); err == nil { + peerConfig.AddressV6 = b + } + } + + return peerConfig +} + +// ToProtoNetworkMap mirrors main's ToSyncResponse, restricted to the +// proto.NetworkMap it produces. SyncResponse-level fields (NetbirdConfig, +// Checks, the deprecated top-level RemotePeers) are omitted — they are not part +// of the equivalence surface. PeerConfig is included because proto.NetworkMap +// carries it, and it is where main's ForceRoutingPeerDNSResolution surfaces. +func ToProtoNetworkMap( + ctx context.Context, + peer *nbpeer.Peer, + nm *NetworkMap, + dnsName string, + settings *types.Settings, + httpConfig *nbconfig.HttpServerConfig, + dnsCache networkmap.DNSConfigCache, + dnsFwdPort int64, +) *proto.NetworkMap { + includeIPv6 := peer.SupportsIPv6() && peer.IPv6.IsValid() + useSourcePrefixes := peer.SupportsSourcePrefixes() + + peerConfig := toPeerConfig(peer, nm.Network, dnsName, settings, httpConfig, nil, nm.EnableSSH, nm.ForceRoutingPeerDNSResolution) + + pm := &proto.NetworkMap{ + Serial: nm.Network.CurrentSerial(), + Routes: ToProtocolRoutes(nm.Routes), + DNSConfig: networkmap.ToProtocolDNSConfig(nm.DNSConfig, dnsCache, dnsFwdPort), + PeerConfig: peerConfig, + } + + remotePeers := make([]*proto.RemotePeerConfig, 0, len(nm.Peers)+len(nm.OfflinePeers)) + remotePeers = AppendRemotePeerConfig(remotePeers, nm.Peers, dnsName, includeIPv6) + pm.RemotePeers = remotePeers + pm.RemotePeersIsEmpty = len(remotePeers) == 0 + + pm.OfflinePeers = AppendRemotePeerConfig(nil, nm.OfflinePeers, dnsName, includeIPv6) + + firewallRules := networkmap.ToProtocolFirewallRules(nm.FirewallRules, includeIPv6, useSourcePrefixes) + pm.FirewallRules = firewallRules + pm.FirewallRulesIsEmpty = len(firewallRules) == 0 + + routesFirewallRules := networkmap.ToProtocolRoutesFirewallRules(nm.RoutesFirewallRules) + pm.RoutesFirewallRules = routesFirewallRules + pm.RoutesFirewallRulesIsEmpty = len(routesFirewallRules) == 0 + + if nm.ForwardingRules != nil { + forwardingRules := make([]*proto.ForwardingRule, 0, len(nm.ForwardingRules)) + for _, rule := range nm.ForwardingRules { + forwardingRules = append(forwardingRules, rule.ToProto()) + } + pm.ForwardingRules = forwardingRules + } + + if nm.AuthorizedUsers != nil { + hashedUsers, machineUsers := networkmap.BuildAuthorizedUsersProto(ctx, nm.AuthorizedUsers) + userIDClaim := auth.DefaultUserIDClaim + if httpConfig != nil && httpConfig.AuthUserIDClaim != "" { + userIDClaim = httpConfig.AuthUserIDClaim + } + pm.SshAuth = &proto.SSHAuth{AuthorizedUsers: hashedUsers, MachineUsers: machineUsers, UserIDClaim: userIDClaim} + } + + return pm +} + +func deriveIssuerFromTokenEndpoint(tokenEndpoint string) string { + if tokenEndpoint == "" { + return "" + } + + u, err := url.Parse(tokenEndpoint) + if err != nil { + return "" + } + + return fmt.Sprintf("%s://%s/", u.Scheme, u.Host) +} diff --git a/shared/management/types/network.go b/shared/management/types/network.go index 1f662ffd9..2e4528b4e 100644 --- a/shared/management/types/network.go +++ b/shared/management/types/network.go @@ -48,6 +48,10 @@ type NetworkMap struct { ForwardingRules []*ForwardingRule AuthorizedUsers map[string]map[string]struct{} EnableSSH bool + // ForceRoutingPeerDNSResolution forces the peer to run/use routing-peer DNS + // resolution regardless of the account-global setting, for reverse-proxy + // domain targets. + ForceRoutingPeerDNSResolution bool } func (nm *NetworkMap) Merge(other *NetworkMap) { @@ -57,6 +61,7 @@ func (nm *NetworkMap) Merge(other *NetworkMap) { nm.FirewallRules = util.MergeUnique(nm.FirewallRules, other.FirewallRules) nm.RoutesFirewallRules = util.MergeUnique(nm.RoutesFirewallRules, other.RoutesFirewallRules) nm.ForwardingRules = util.MergeUnique(nm.ForwardingRules, other.ForwardingRules) + nm.ForceRoutingPeerDNSResolution = nm.ForceRoutingPeerDNSResolution || other.ForceRoutingPeerDNSResolution } func mergeUniquePeersByID(peers1, peers2 []*nmdata.Peer) []*nmdata.Peer { diff --git a/shared/management/types/networkmap_components.go b/shared/management/types/networkmap_components.go index e28dbd095..948af6169 100644 --- a/shared/management/types/networkmap_components.go +++ b/shared/management/types/networkmap_components.go @@ -53,8 +53,14 @@ type NetworkMapComponents struct { // Same role as NetworkXIDToPublicID, used for PostureFailedPeers keys and // policy SourcePostureChecks references. PostureCheckXIDToPublicID map[string]string - routesByPeerOnce sync.Once - routesByPeerIdx map[string][]routeIndexEntry + + // ForceRoutingPeerDNSResolution forces the peer to run/use routing-peer DNS + // resolution regardless of the account-global setting, for reverse-proxy + // domain targets. + ForceRoutingPeerDNSResolution bool + + routesByPeerOnce sync.Once + routesByPeerIdx map[string][]routeIndexEntry // true when returning an empty-like map (returned instead of nil) empty bool @@ -192,6 +198,8 @@ func (c *NetworkMapComponents) Calculate(ctx context.Context) *NetworkMap { RoutesFirewallRules: append(networkResourcesFirewallRules, routesFirewallRules...), AuthorizedUsers: authorizedUsers, EnableSSH: sshEnabled, + + ForceRoutingPeerDNSResolution: c.ForceRoutingPeerDNSResolution, } }