mirror of
https://github.com/netbirdio/netbird.git
synced 2026-08-29 11:01:29 +02:00
Merge branch 'main' into fix/remove-math-rand
This commit is contained in:
103
shared/management/types/component_types.go
Normal file
103
shared/management/types/component_types.go
Normal file
@@ -0,0 +1,103 @@
|
||||
package types
|
||||
|
||||
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
|
||||
}
|
||||
16
shared/management/types/dns_settings.go
Normal file
16
shared/management/types/dns_settings.go
Normal file
@@ -0,0 +1,16 @@
|
||||
package types
|
||||
|
||||
// DNSSettings defines dns settings at the account level
|
||||
type DNSSettings struct {
|
||||
// DisabledManagementGroups groups whose DNS management is disabled
|
||||
DisabledManagementGroups []string `gorm:"serializer:json"`
|
||||
}
|
||||
|
||||
// Copy returns a copy of the DNS settings
|
||||
func (d DNSSettings) Copy() DNSSettings {
|
||||
settings := DNSSettings{
|
||||
DisabledManagementGroups: make([]string, len(d.DisabledManagementGroups)),
|
||||
}
|
||||
copy(settings.DisabledManagementGroups, d.DisabledManagementGroups)
|
||||
return settings
|
||||
}
|
||||
129
shared/management/types/firewall_helpers.go
Normal file
129
shared/management/types/firewall_helpers.go
Normal file
@@ -0,0 +1,129 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
|
||||
"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 := version.MeetsMinVersion(firewallRuleMinNativeSSHVer, peerVer)
|
||||
features.nativeSSH = err == nil && meetMinVer
|
||||
|
||||
if features.nativeSSH {
|
||||
features.portRanges = true
|
||||
} else {
|
||||
meetMinVer, err = version.MeetsMinVersion(firewallRuleMinPortRangesVer, peerVer)
|
||||
features.portRanges = err == nil && meetMinVer
|
||||
}
|
||||
|
||||
return features
|
||||
}
|
||||
181
shared/management/types/firewall_rule.go
Normal file
181
shared/management/types/firewall_rule.go
Normal file
@@ -0,0 +1,181 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
nbroute "github.com/netbirdio/netbird/route"
|
||||
)
|
||||
|
||||
const (
|
||||
FirewallRuleDirectionIN = 0
|
||||
FirewallRuleDirectionOUT = 1
|
||||
)
|
||||
|
||||
// FirewallRule is a rule of the firewall.
|
||||
type FirewallRule struct {
|
||||
// PolicyID is the ID of the policy this rule is derived from
|
||||
PolicyID string
|
||||
|
||||
// PeerIP of the peer
|
||||
PeerIP string
|
||||
|
||||
// Direction of the traffic
|
||||
Direction int
|
||||
|
||||
// Action of the traffic
|
||||
Action string
|
||||
|
||||
// Protocol of the traffic
|
||||
Protocol string
|
||||
|
||||
// Port of the traffic
|
||||
Port string
|
||||
|
||||
// PortRange represents the range of ports for a firewall rule
|
||||
PortRange RulePortRange
|
||||
}
|
||||
|
||||
// Equal checks if two firewall rules are equal.
|
||||
func (r *FirewallRule) Equal(other *FirewallRule) bool {
|
||||
return reflect.DeepEqual(r, other)
|
||||
}
|
||||
|
||||
// GenerateRouteFirewallRules generates a list of firewall rules for a given route.
|
||||
// For static routes, source ranges match the destination family (v4 or v6).
|
||||
// For dynamic routes (domain-based), separate v4 and v6 rules are generated
|
||||
// so the routing peer's forwarding chain allows both address families.
|
||||
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
|
||||
}
|
||||
|
||||
// splitPeerSourcesByFamily separates peer IPs into v4 (/32) and v6 (/128) source ranges.
|
||||
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
|
||||
}
|
||||
|
||||
// generateRulesForPeer generates rules for a given peer based on ports and port ranges.
|
||||
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
|
||||
}
|
||||
|
||||
// generateRulesWithPorts generates rules when specific ports are provided.
|
||||
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
|
||||
}
|
||||
|
||||
// generateRuleIDBase generates the base rule ID for checking duplicates.
|
||||
func generateRuleIDBase(rule *PolicyRule, baseRule RouteFirewallRule) string {
|
||||
return rule.ID + strings.Join(baseRule.SourceRanges, ",") + strconv.Itoa(FirewallRuleDirectionIN) + baseRule.Protocol + baseRule.Action
|
||||
}
|
||||
196
shared/management/types/firewall_rule_test.go
Normal file
196
shared/management/types/firewall_rule_test.go
Normal file
@@ -0,0 +1,196 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/route"
|
||||
"github.com/netbirdio/netbird/shared/management/domain"
|
||||
)
|
||||
|
||||
func TestSplitPeerSourcesByFamily(t *testing.T) {
|
||||
peers := []*ComponentPeer{
|
||||
{
|
||||
IP: netip.MustParseAddr("100.64.0.1"),
|
||||
IPv6: netip.MustParseAddr("fd00::1"),
|
||||
},
|
||||
{
|
||||
IP: netip.MustParseAddr("100.64.0.2"),
|
||||
},
|
||||
{
|
||||
IP: netip.MustParseAddr("100.64.0.3"),
|
||||
IPv6: netip.MustParseAddr("fd00::3"),
|
||||
},
|
||||
nil,
|
||||
}
|
||||
|
||||
v4, v6 := splitPeerSourcesByFamily(peers)
|
||||
|
||||
assert.Equal(t, []string{"100.64.0.1/32", "100.64.0.2/32", "100.64.0.3/32"}, v4)
|
||||
assert.Equal(t, []string{"fd00::1/128", "fd00::3/128"}, v6)
|
||||
}
|
||||
|
||||
func TestGenerateRouteFirewallRules_V4Route(t *testing.T) {
|
||||
peers := []*ComponentPeer{
|
||||
{
|
||||
IP: netip.MustParseAddr("100.64.0.1"),
|
||||
IPv6: netip.MustParseAddr("fd00::1"),
|
||||
},
|
||||
{
|
||||
IP: netip.MustParseAddr("100.64.0.2"),
|
||||
},
|
||||
}
|
||||
|
||||
r := &route.Route{
|
||||
ID: "route1",
|
||||
Network: netip.MustParsePrefix("10.0.0.0/24"),
|
||||
}
|
||||
rule := &PolicyRule{
|
||||
PolicyID: "policy1",
|
||||
ID: "rule1",
|
||||
Action: PolicyTrafficActionAccept,
|
||||
Protocol: PolicyRuleProtocolALL,
|
||||
}
|
||||
|
||||
rules := GenerateRouteFirewallRules(context.Background(), r, rule, peers, FirewallRuleDirectionIN, true)
|
||||
|
||||
require.Len(t, rules, 1)
|
||||
assert.Equal(t, []string{"100.64.0.1/32", "100.64.0.2/32"}, rules[0].SourceRanges, "v4 route should only have v4 sources")
|
||||
assert.Equal(t, "10.0.0.0/24", rules[0].Destination)
|
||||
}
|
||||
|
||||
func TestGenerateRouteFirewallRules_V6Route(t *testing.T) {
|
||||
peers := []*ComponentPeer{
|
||||
{
|
||||
IP: netip.MustParseAddr("100.64.0.1"),
|
||||
IPv6: netip.MustParseAddr("fd00::1"),
|
||||
},
|
||||
{
|
||||
IP: netip.MustParseAddr("100.64.0.2"),
|
||||
},
|
||||
}
|
||||
|
||||
r := &route.Route{
|
||||
ID: "route1",
|
||||
Network: netip.MustParsePrefix("2001:db8::/32"),
|
||||
}
|
||||
rule := &PolicyRule{
|
||||
PolicyID: "policy1",
|
||||
ID: "rule1",
|
||||
Action: PolicyTrafficActionAccept,
|
||||
Protocol: PolicyRuleProtocolALL,
|
||||
}
|
||||
|
||||
rules := GenerateRouteFirewallRules(context.Background(), r, rule, peers, FirewallRuleDirectionIN, true)
|
||||
|
||||
require.Len(t, rules, 1)
|
||||
assert.Equal(t, []string{"fd00::1/128"}, rules[0].SourceRanges, "v6 route should only have v6 sources")
|
||||
}
|
||||
|
||||
func TestGenerateRouteFirewallRules_DynamicRoute_DualStack(t *testing.T) {
|
||||
peers := []*ComponentPeer{
|
||||
{
|
||||
IP: netip.MustParseAddr("100.64.0.1"),
|
||||
IPv6: netip.MustParseAddr("fd00::1"),
|
||||
},
|
||||
{
|
||||
IP: netip.MustParseAddr("100.64.0.2"),
|
||||
},
|
||||
}
|
||||
|
||||
r := &route.Route{
|
||||
ID: "route1",
|
||||
NetworkType: route.DomainNetwork,
|
||||
Domains: domain.List{"example.com"},
|
||||
}
|
||||
rule := &PolicyRule{
|
||||
PolicyID: "policy1",
|
||||
ID: "rule1",
|
||||
Action: PolicyTrafficActionAccept,
|
||||
Protocol: PolicyRuleProtocolALL,
|
||||
}
|
||||
|
||||
rules := GenerateRouteFirewallRules(context.Background(), r, rule, peers, FirewallRuleDirectionIN, true)
|
||||
|
||||
require.Len(t, rules, 2, "dynamic route should produce both v4 and v6 rules")
|
||||
assert.Equal(t, []string{"100.64.0.1/32", "100.64.0.2/32"}, rules[0].SourceRanges)
|
||||
assert.Equal(t, []string{"fd00::1/128"}, rules[1].SourceRanges)
|
||||
assert.Equal(t, rules[0].Domains, rules[1].Domains)
|
||||
assert.True(t, rules[0].IsDynamic)
|
||||
assert.True(t, rules[1].IsDynamic)
|
||||
}
|
||||
|
||||
func TestGenerateRouteFirewallRules_DynamicRoute_NoV6Peers(t *testing.T) {
|
||||
peers := []*ComponentPeer{
|
||||
{IP: netip.MustParseAddr("100.64.0.1")},
|
||||
{IP: netip.MustParseAddr("100.64.0.2")},
|
||||
}
|
||||
|
||||
r := &route.Route{
|
||||
ID: "route1",
|
||||
NetworkType: route.DomainNetwork,
|
||||
Domains: domain.List{"example.com"},
|
||||
}
|
||||
rule := &PolicyRule{
|
||||
PolicyID: "policy1",
|
||||
ID: "rule1",
|
||||
Action: PolicyTrafficActionAccept,
|
||||
Protocol: PolicyRuleProtocolALL,
|
||||
}
|
||||
|
||||
rules := GenerateRouteFirewallRules(context.Background(), r, rule, peers, FirewallRuleDirectionIN, true)
|
||||
|
||||
require.Len(t, rules, 1, "no v6 peers means only v4 rule")
|
||||
assert.Equal(t, []string{"100.64.0.1/32", "100.64.0.2/32"}, rules[0].SourceRanges)
|
||||
}
|
||||
|
||||
func TestGenerateRouteFirewallRules_IncludeIPv6False(t *testing.T) {
|
||||
peers := []*ComponentPeer{
|
||||
{
|
||||
IP: netip.MustParseAddr("100.64.0.1"),
|
||||
IPv6: netip.MustParseAddr("fd00::1"),
|
||||
},
|
||||
{
|
||||
IP: netip.MustParseAddr("100.64.0.2"),
|
||||
IPv6: netip.MustParseAddr("fd00::2"),
|
||||
},
|
||||
}
|
||||
|
||||
t.Run("v6 route excluded", func(t *testing.T) {
|
||||
r := &route.Route{
|
||||
ID: "route1",
|
||||
Network: netip.MustParsePrefix("2001:db8::/32"),
|
||||
}
|
||||
rule := &PolicyRule{
|
||||
PolicyID: "policy1",
|
||||
ID: "rule1",
|
||||
Action: PolicyTrafficActionAccept,
|
||||
Protocol: PolicyRuleProtocolALL,
|
||||
}
|
||||
|
||||
rules := GenerateRouteFirewallRules(context.Background(), r, rule, peers, FirewallRuleDirectionIN, false)
|
||||
assert.Empty(t, rules, "v6 route should produce no rules when includeIPv6 is false")
|
||||
})
|
||||
|
||||
t.Run("dynamic route only v4", func(t *testing.T) {
|
||||
r := &route.Route{
|
||||
ID: "route1",
|
||||
NetworkType: route.DomainNetwork,
|
||||
Domains: domain.List{"example.com"},
|
||||
}
|
||||
rule := &PolicyRule{
|
||||
PolicyID: "policy1",
|
||||
ID: "rule1",
|
||||
Action: PolicyTrafficActionAccept,
|
||||
Protocol: PolicyRuleProtocolALL,
|
||||
}
|
||||
|
||||
rules := GenerateRouteFirewallRules(context.Background(), r, rule, peers, FirewallRuleDirectionIN, false)
|
||||
require.Len(t, rules, 1, "dynamic route with includeIPv6=false should produce only v4 rule")
|
||||
assert.Equal(t, []string{"100.64.0.1/32", "100.64.0.2/32"}, rules[0].SourceRanges)
|
||||
})
|
||||
}
|
||||
407
shared/management/types/network.go
Normal file
407
shared/management/types/network.go
Normal file
@@ -0,0 +1,407 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"sync"
|
||||
|
||||
"github.com/c-robinson/iplib"
|
||||
"github.com/rs/xid"
|
||||
"golang.org/x/exp/maps"
|
||||
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
"github.com/netbirdio/netbird/route"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
const (
|
||||
// SubnetSize is a size of the subnet of the global network, e.g. 100.77.0.0/16
|
||||
SubnetSize = 16
|
||||
// NetSize is a global network size 100.64.0.0/10
|
||||
NetSize = 10
|
||||
|
||||
// AllowedIPsFormat generates Wireguard AllowedIPs format (e.g. 100.64.30.1/32)
|
||||
AllowedIPsFormat = "%s/32"
|
||||
// AllowedIPsV6Format generates AllowedIPs format for v6 (e.g. fd12:3456:7890::1/128)
|
||||
AllowedIPsV6Format = "%s/128"
|
||||
|
||||
// IPv6SubnetSize is the prefix length of per-account IPv6 subnets.
|
||||
// Each account gets a /64 from its unique /48 ULA prefix.
|
||||
IPv6SubnetSize = 64
|
||||
)
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
func (nm *NetworkMap) Merge(other *NetworkMap) {
|
||||
nm.Peers = mergeUniquePeersByID(nm.Peers, other.Peers)
|
||||
nm.Routes = mergeUnique(nm.Routes, other.Routes)
|
||||
nm.OfflinePeers = mergeUniquePeersByID(nm.OfflinePeers, other.OfflinePeers)
|
||||
nm.FirewallRules = mergeUnique(nm.FirewallRules, other.FirewallRules)
|
||||
nm.RoutesFirewallRules = mergeUnique(nm.RoutesFirewallRules, other.RoutesFirewallRules)
|
||||
nm.ForwardingRules = mergeUnique(nm.ForwardingRules, other.ForwardingRules)
|
||||
nm.ForceRoutingPeerDNSResolution = nm.ForceRoutingPeerDNSResolution || other.ForceRoutingPeerDNSResolution
|
||||
}
|
||||
|
||||
type comparableObject[T any] interface {
|
||||
Equal(other T) bool
|
||||
}
|
||||
|
||||
func mergeUnique[T comparableObject[T]](arr1, arr2 []T) []T {
|
||||
var result []T
|
||||
|
||||
for _, item := range arr1 {
|
||||
if !containsEqual(result, item) {
|
||||
result = append(result, item)
|
||||
}
|
||||
}
|
||||
|
||||
for _, item := range arr2 {
|
||||
if !containsEqual(result, item) {
|
||||
result = append(result, item)
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
func containsEqual[T comparableObject[T]](slice []T, element T) bool {
|
||||
for _, item := range slice {
|
||||
if item.Equal(element) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func mergeUniquePeersByID(peers1, peers2 []*ComponentPeer) []*ComponentPeer {
|
||||
result := make(map[string]*ComponentPeer)
|
||||
for _, peer := range peers1 {
|
||||
result[peer.ID] = peer
|
||||
}
|
||||
for _, peer := range peers2 {
|
||||
if _, ok := result[peer.ID]; !ok {
|
||||
result[peer.ID] = peer
|
||||
}
|
||||
}
|
||||
|
||||
return maps.Values(result)
|
||||
}
|
||||
|
||||
type ForwardingRule struct {
|
||||
RuleProtocol string
|
||||
DestinationPorts RulePortRange
|
||||
TranslatedAddress net.IP
|
||||
TranslatedPorts RulePortRange
|
||||
}
|
||||
|
||||
func (f *ForwardingRule) ToProto() *proto.ForwardingRule {
|
||||
var protocol proto.RuleProtocol
|
||||
switch f.RuleProtocol {
|
||||
case "icmp":
|
||||
protocol = proto.RuleProtocol_ICMP
|
||||
case "tcp":
|
||||
protocol = proto.RuleProtocol_TCP
|
||||
case "udp":
|
||||
protocol = proto.RuleProtocol_UDP
|
||||
case "all":
|
||||
protocol = proto.RuleProtocol_ALL
|
||||
default:
|
||||
protocol = proto.RuleProtocol_UNKNOWN
|
||||
}
|
||||
return &proto.ForwardingRule{
|
||||
Protocol: protocol,
|
||||
DestinationPort: f.DestinationPorts.ToProto(),
|
||||
TranslatedAddress: ipToBytes(f.TranslatedAddress),
|
||||
TranslatedPort: f.TranslatedPorts.ToProto(),
|
||||
}
|
||||
}
|
||||
|
||||
func (f *ForwardingRule) Equal(other *ForwardingRule) bool {
|
||||
return f.RuleProtocol == other.RuleProtocol &&
|
||||
f.DestinationPorts.Equal(&other.DestinationPorts) &&
|
||||
f.TranslatedAddress.Equal(other.TranslatedAddress) &&
|
||||
f.TranslatedPorts.Equal(&other.TranslatedPorts)
|
||||
}
|
||||
|
||||
func ipToBytes(ip net.IP) []byte {
|
||||
if ip4 := ip.To4(); ip4 != nil {
|
||||
return ip4
|
||||
}
|
||||
return ip.To16()
|
||||
}
|
||||
|
||||
type Network struct {
|
||||
Identifier string `json:"id"`
|
||||
Net net.IPNet `gorm:"serializer:json"`
|
||||
// NetV6 is the IPv6 ULA subnet for this account's overlay. Empty if not yet allocated.
|
||||
NetV6 net.IPNet `gorm:"serializer:json"`
|
||||
Dns string
|
||||
// Serial is an ID that increments by 1 when any change to the network happened (e.g. new peer has been added).
|
||||
// Used to synchronize state to the client apps.
|
||||
Serial uint64
|
||||
|
||||
Mu sync.Mutex `json:"-" gorm:"-"`
|
||||
}
|
||||
|
||||
// NewNetwork creates a new Network initializing it with a Serial=0
|
||||
// It takes a random /16 subnet from 100.64.0.0/10 (64 different subnets)
|
||||
// and a random /64 subnet from fd00:4e42::/32 for IPv6.
|
||||
func NewNetwork() *Network {
|
||||
n := iplib.NewNet4(net.ParseIP("100.64.0.0"), NetSize)
|
||||
sub, _ := n.Subnet(SubnetSize)
|
||||
|
||||
intn := util.RandIntn(len(sub))
|
||||
|
||||
return &Network{
|
||||
Identifier: xid.New().String(),
|
||||
Net: sub[intn].IPNet,
|
||||
NetV6: AllocateIPv6Subnet(),
|
||||
Dns: "",
|
||||
Serial: 0,
|
||||
}
|
||||
}
|
||||
|
||||
// AllocateIPv6Subnet generates a random RFC 4193 ULA /64 prefix.
|
||||
// The format follows RFC 4193 section 3.1: fd + 40-bit Global ID + 16-bit Subnet ID.
|
||||
// The Global ID and Subnet ID are randomized (simplified from the SHA-1 algorithm
|
||||
// in section 3.2.2), giving 2^56 possible /64 subnets across all accounts.
|
||||
func AllocateIPv6Subnet() net.IPNet {
|
||||
ip := make(net.IP, 16)
|
||||
ip[0] = 0xfd
|
||||
// Bytes 1-5: 40-bit random Global ID, bytes 6-7: 16-bit random Subnet ID
|
||||
if _, err := rand.Read(ip[1:8]); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
return net.IPNet{
|
||||
IP: ip,
|
||||
Mask: net.CIDRMask(IPv6SubnetSize, 128),
|
||||
}
|
||||
}
|
||||
|
||||
// IncSerial increments Serial by 1 reflecting that the network state has been changed
|
||||
func (n *Network) IncSerial() {
|
||||
n.Mu.Lock()
|
||||
defer n.Mu.Unlock()
|
||||
n.Serial++
|
||||
}
|
||||
|
||||
// CurrentSerial returns the Network.Serial of the network (latest state id)
|
||||
func (n *Network) CurrentSerial() uint64 {
|
||||
n.Mu.Lock()
|
||||
defer n.Mu.Unlock()
|
||||
return n.Serial
|
||||
}
|
||||
|
||||
func (n *Network) Copy() *Network {
|
||||
n.Mu.Lock()
|
||||
defer n.Mu.Unlock()
|
||||
return &Network{
|
||||
Identifier: n.Identifier,
|
||||
Net: n.Net,
|
||||
NetV6: n.NetV6,
|
||||
Dns: n.Dns,
|
||||
Serial: n.Serial,
|
||||
}
|
||||
}
|
||||
|
||||
// validateIPv4Prefix ensures the prefix is an IPv4 network with assignable host addresses.
|
||||
func validateIPv4Prefix(prefix netip.Prefix) error {
|
||||
if !prefix.IsValid() || !prefix.Addr().Is4() || prefix.Bits() < 1 || prefix.Bits() >= 31 {
|
||||
return fmt.Errorf("invalid IPv4 subnet: %s", prefix.String())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// AllocatePeerIP picks an available IP from a netip.Prefix.
|
||||
// This method considers already taken IPs and reuses IPs if there are gaps in takenIps.
|
||||
// E.g. if prefix=100.30.0.0/16 and takenIps=[100.30.0.1, 100.30.0.4] then the result would be 100.30.0.2 or 100.30.0.3.
|
||||
func AllocatePeerIP(prefix netip.Prefix, takenIps []netip.Addr) (netip.Addr, error) {
|
||||
if err := validateIPv4Prefix(prefix); err != nil {
|
||||
return netip.Addr{}, err
|
||||
}
|
||||
|
||||
b := prefix.Masked().Addr().As4()
|
||||
baseIP := binary.BigEndian.Uint32(b[:])
|
||||
hostBits := 32 - prefix.Bits()
|
||||
totalIPs := uint32(1 << hostBits)
|
||||
|
||||
taken := make(map[uint32]struct{}, len(takenIps)+1)
|
||||
taken[baseIP] = struct{}{} // reserve network IP
|
||||
taken[baseIP+totalIPs-1] = struct{}{} // reserve broadcast IP
|
||||
|
||||
for _, ip := range takenIps {
|
||||
if !ip.Is4() {
|
||||
continue
|
||||
}
|
||||
ab := ip.As4()
|
||||
taken[binary.BigEndian.Uint32(ab[:])] = struct{}{}
|
||||
}
|
||||
|
||||
maxAttempts := (int(totalIPs) - len(taken)) / 100
|
||||
|
||||
for i := 0; i < maxAttempts; i++ {
|
||||
offset := uint32(util.RandIntn(int(totalIPs-2))) + 1
|
||||
candidate := baseIP + offset
|
||||
if _, exists := taken[candidate]; !exists {
|
||||
return uint32ToIP(candidate), nil
|
||||
}
|
||||
}
|
||||
|
||||
for offset := uint32(1); offset < totalIPs-1; offset++ {
|
||||
candidate := baseIP + offset
|
||||
if _, exists := taken[candidate]; !exists {
|
||||
return uint32ToIP(candidate), nil
|
||||
}
|
||||
}
|
||||
|
||||
return netip.Addr{}, status.Errorf(status.PreconditionFailed, "network %s is out of IPs", prefix.String())
|
||||
}
|
||||
|
||||
// AllocateRandomPeerIP picks a random available IP from a netip.Prefix.
|
||||
func AllocateRandomPeerIP(prefix netip.Prefix) (netip.Addr, error) {
|
||||
if err := validateIPv4Prefix(prefix); err != nil {
|
||||
return netip.Addr{}, err
|
||||
}
|
||||
|
||||
b := prefix.Masked().Addr().As4()
|
||||
baseIP := binary.BigEndian.Uint32(b[:])
|
||||
hostBits := 32 - prefix.Bits()
|
||||
totalIPs := uint32(1 << hostBits)
|
||||
|
||||
offset := uint32(util.RandIntn(int(totalIPs-2))) + 1
|
||||
|
||||
candidate := baseIP + offset
|
||||
return uint32ToIP(candidate), nil
|
||||
}
|
||||
|
||||
// AllocateRandomPeerIPv6 picks a random host address within the given IPv6 prefix.
|
||||
// Only the host bits (after the prefix length) are randomized.
|
||||
func AllocateRandomPeerIPv6(prefix netip.Prefix) (netip.Addr, error) {
|
||||
ones := prefix.Bits()
|
||||
if ones == 0 || ones > 126 || !prefix.Addr().Is6() {
|
||||
return netip.Addr{}, fmt.Errorf("invalid IPv6 subnet: %s", prefix.String())
|
||||
}
|
||||
|
||||
ip := prefix.Addr().As16()
|
||||
|
||||
// Determine which byte the host bits start in
|
||||
firstHostByte := ones / 8
|
||||
// If the prefix doesn't end on a byte boundary, handle the partial byte
|
||||
partialBits := ones % 8
|
||||
|
||||
var rnd [16]byte
|
||||
if _, err := rand.Read(rnd[firstHostByte:]); err != nil {
|
||||
return netip.Addr{}, err
|
||||
}
|
||||
|
||||
if partialBits > 0 {
|
||||
// Keep the network bits in the partial byte, randomize the rest
|
||||
hostMask := byte(0xff >> partialBits)
|
||||
ip[firstHostByte] = (ip[firstHostByte] & ^hostMask) | (rnd[firstHostByte] & hostMask)
|
||||
firstHostByte++
|
||||
}
|
||||
|
||||
// Randomize remaining full host bytes
|
||||
for i := firstHostByte; i < 16; i++ {
|
||||
ip[i] = rnd[i]
|
||||
}
|
||||
|
||||
// Avoid all-zeros and all-ones host parts by checking only host bits.
|
||||
if isHostAllZeroOrOnes(ip[:], ones) {
|
||||
ip = prefix.Masked().Addr().As16()
|
||||
ip[15] |= 0x01
|
||||
}
|
||||
|
||||
return netip.AddrFrom16(ip).Unmap(), nil
|
||||
}
|
||||
|
||||
// isHostAllZeroOrOnes checks whether all host bits (after prefixLen) are zero or all ones.
|
||||
func isHostAllZeroOrOnes(ip []byte, prefixLen int) bool {
|
||||
hostStart := prefixLen / 8
|
||||
partialBits := prefixLen % 8
|
||||
|
||||
hostSlice := slices.Clone(ip[hostStart:])
|
||||
if partialBits > 0 {
|
||||
hostSlice[0] &= 0xff >> partialBits
|
||||
}
|
||||
|
||||
allZero := !slices.ContainsFunc(hostSlice, func(v byte) bool { return v != 0 })
|
||||
if allZero {
|
||||
return true
|
||||
}
|
||||
|
||||
// Build the all-ones mask for host bits
|
||||
onesMask := make([]byte, len(hostSlice))
|
||||
for i := range onesMask {
|
||||
onesMask[i] = 0xff
|
||||
}
|
||||
if partialBits > 0 {
|
||||
onesMask[0] = 0xff >> partialBits
|
||||
}
|
||||
|
||||
return slices.Equal(hostSlice, onesMask)
|
||||
}
|
||||
|
||||
func uint32ToIP(n uint32) netip.Addr {
|
||||
var b [4]byte
|
||||
binary.BigEndian.PutUint32(b[:], n)
|
||||
return netip.AddrFrom4(b)
|
||||
}
|
||||
|
||||
// generateIPs generates a list of all possible IPs of the given network excluding IPs specified in the exclusion list
|
||||
func generateIPs(ipNet *net.IPNet, exclusions map[string]struct{}) ([]net.IP, int) {
|
||||
|
||||
var ips []net.IP
|
||||
for ip := ipNet.IP.Mask(ipNet.Mask); ipNet.Contains(ip); incIP(ip) {
|
||||
if _, ok := exclusions[ip.String()]; !ok && ip[3] != 0 {
|
||||
ips = append(ips, copyIP(ip))
|
||||
}
|
||||
}
|
||||
|
||||
// remove network address, broadcast and Fake DNS resolver address
|
||||
lenIPs := len(ips)
|
||||
switch {
|
||||
case lenIPs < 2:
|
||||
return ips, lenIPs
|
||||
case lenIPs < 3:
|
||||
return ips[1 : len(ips)-1], lenIPs - 2
|
||||
default:
|
||||
return ips[1 : len(ips)-2], lenIPs - 3
|
||||
}
|
||||
}
|
||||
|
||||
func copyIP(ip net.IP) net.IP {
|
||||
dup := make(net.IP, len(ip))
|
||||
copy(dup, ip)
|
||||
return dup
|
||||
}
|
||||
|
||||
func incIP(ip net.IP) {
|
||||
for j := len(ip) - 1; j >= 0; j-- {
|
||||
ip[j]++
|
||||
if ip[j] > 0 {
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
41
shared/management/types/network_merge_test.go
Normal file
41
shared/management/types/network_merge_test.go
Normal file
@@ -0,0 +1,41 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
type testObject struct {
|
||||
value int
|
||||
}
|
||||
|
||||
func (t testObject) Equal(other testObject) bool {
|
||||
return t.value == other.value
|
||||
}
|
||||
|
||||
func Test_MergeUniqueArraysWithoutDuplicates(t *testing.T) {
|
||||
arr1 := []testObject{{value: 1}, {value: 2}}
|
||||
arr2 := []testObject{{value: 2}, {value: 3}}
|
||||
result := mergeUnique(arr1, arr2)
|
||||
assert.Len(t, result, 3)
|
||||
assert.Contains(t, result, testObject{value: 1})
|
||||
assert.Contains(t, result, testObject{value: 2})
|
||||
assert.Contains(t, result, testObject{value: 3})
|
||||
}
|
||||
|
||||
func Test_MergeUniqueHandlesEmptyArrays(t *testing.T) {
|
||||
arr1 := []testObject{}
|
||||
arr2 := []testObject{}
|
||||
result := mergeUnique(arr1, arr2)
|
||||
assert.Empty(t, result)
|
||||
}
|
||||
|
||||
func Test_MergeUniqueHandlesOneEmptyArray(t *testing.T) {
|
||||
arr1 := []testObject{{value: 1}, {value: 2}}
|
||||
arr2 := []testObject{}
|
||||
result := mergeUnique(arr1, arr2)
|
||||
assert.Len(t, result, 2)
|
||||
assert.Contains(t, result, testObject{value: 1})
|
||||
assert.Contains(t, result, testObject{value: 2})
|
||||
}
|
||||
292
shared/management/types/network_test.go
Normal file
292
shared/management/types/network_test.go
Normal file
@@ -0,0 +1,292 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"net"
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestNewNetwork(t *testing.T) {
|
||||
network := NewNetwork()
|
||||
|
||||
// generated net should be a subnet of a larger 100.64.0.0/10 net
|
||||
ipNet := net.IPNet{IP: net.ParseIP("100.64.0.0"), Mask: net.IPMask{255, 192, 0, 0}}
|
||||
assert.Equal(t, ipNet.Contains(network.Net.IP), true)
|
||||
}
|
||||
|
||||
func TestAllocatePeerIP(t *testing.T) {
|
||||
prefix := netip.MustParsePrefix("100.64.0.0/24")
|
||||
var ips []netip.Addr
|
||||
for i := 0; i < 252; i++ {
|
||||
ip, err := AllocatePeerIP(prefix, ips)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ips = append(ips, ip)
|
||||
}
|
||||
|
||||
assert.Len(t, ips, 252)
|
||||
|
||||
uniq := make(map[string]struct{})
|
||||
for _, ip := range ips {
|
||||
if _, ok := uniq[ip.String()]; !ok {
|
||||
uniq[ip.String()] = struct{}{}
|
||||
} else {
|
||||
t.Errorf("found duplicate IP %s", ip.String())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAllocatePeerIPSmallSubnet(t *testing.T) {
|
||||
// Test /27 network (10.0.0.0/27) - should only have 30 usable IPs (10.0.0.1 to 10.0.0.30)
|
||||
prefix := netip.MustParsePrefix("10.0.0.0/27")
|
||||
var ips []netip.Addr
|
||||
|
||||
// Allocate all available IPs in the /27 network
|
||||
for i := 0; i < 30; i++ {
|
||||
ip, err := AllocatePeerIP(prefix, ips)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// Verify IP is within the correct range
|
||||
if !prefix.Contains(ip) {
|
||||
t.Errorf("allocated IP %s is not within network %s", ip.String(), prefix.String())
|
||||
}
|
||||
|
||||
ips = append(ips, ip)
|
||||
}
|
||||
|
||||
assert.Len(t, ips, 30)
|
||||
|
||||
// Verify all IPs are unique
|
||||
uniq := make(map[string]struct{})
|
||||
for _, ip := range ips {
|
||||
if _, ok := uniq[ip.String()]; !ok {
|
||||
uniq[ip.String()] = struct{}{}
|
||||
} else {
|
||||
t.Errorf("found duplicate IP %s", ip.String())
|
||||
}
|
||||
}
|
||||
|
||||
// Try to allocate one more IP - should fail as network is full
|
||||
_, err := AllocatePeerIP(prefix, ips)
|
||||
if err == nil {
|
||||
t.Error("expected error when network is full, but got none")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAllocatePeerIPVariousCIDRs(t *testing.T) {
|
||||
testCases := []struct {
|
||||
name string
|
||||
cidr string
|
||||
expectedUsable int
|
||||
}{
|
||||
{"/30 network", "192.168.1.0/30", 2}, // 4 total - 2 reserved = 2 usable
|
||||
{"/29 network", "192.168.1.0/29", 6}, // 8 total - 2 reserved = 6 usable
|
||||
{"/28 network", "192.168.1.0/28", 14}, // 16 total - 2 reserved = 14 usable
|
||||
{"/27 network", "192.168.1.0/27", 30}, // 32 total - 2 reserved = 30 usable
|
||||
{"/26 network", "192.168.1.0/26", 62}, // 64 total - 2 reserved = 62 usable
|
||||
{"/25 network", "192.168.1.0/25", 126}, // 128 total - 2 reserved = 126 usable
|
||||
{"/16 network", "10.0.0.0/16", 65534}, // 65536 total - 2 reserved = 65534 usable
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
prefix, err := netip.ParsePrefix(tc.cidr)
|
||||
require.NoError(t, err)
|
||||
prefix = prefix.Masked()
|
||||
|
||||
var ips []netip.Addr
|
||||
|
||||
// For larger networks, test only a subset to avoid long test runs
|
||||
testCount := tc.expectedUsable
|
||||
if testCount > 1000 {
|
||||
testCount = 1000
|
||||
}
|
||||
|
||||
// Allocate IPs and verify they're within the correct range
|
||||
for i := 0; i < testCount; i++ {
|
||||
ip, err := AllocatePeerIP(prefix, ips)
|
||||
require.NoError(t, err, "failed to allocate IP %d", i)
|
||||
|
||||
// Verify IP is within the correct range
|
||||
assert.True(t, prefix.Contains(ip), "allocated IP %s is not within network %s", ip.String(), prefix.String())
|
||||
|
||||
// Verify IP is not network or broadcast address
|
||||
networkAddr := prefix.Masked().Addr()
|
||||
hostBits := 32 - prefix.Bits()
|
||||
b := networkAddr.As4()
|
||||
baseIP := binary.BigEndian.Uint32(b[:])
|
||||
broadcastIP := uint32ToIP(baseIP + (1 << hostBits) - 1)
|
||||
|
||||
assert.NotEqual(t, networkAddr, ip, "allocated network address %s", ip.String())
|
||||
assert.NotEqual(t, broadcastIP, ip, "allocated broadcast address %s", ip.String())
|
||||
|
||||
ips = append(ips, ip)
|
||||
}
|
||||
|
||||
assert.Len(t, ips, testCount)
|
||||
|
||||
// Verify all IPs are unique
|
||||
uniq := make(map[string]struct{})
|
||||
for _, ip := range ips {
|
||||
ipStr := ip.String()
|
||||
assert.NotContains(t, uniq, ipStr, "found duplicate IP %s", ipStr)
|
||||
uniq[ipStr] = struct{}{}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAllocateIPv4InvalidPrefixes(t *testing.T) {
|
||||
prefixes := []netip.Prefix{
|
||||
{},
|
||||
netip.MustParsePrefix("0.0.0.0/0"),
|
||||
netip.MustParsePrefix("192.168.1.0/31"),
|
||||
netip.MustParsePrefix("192.168.1.1/32"),
|
||||
netip.MustParsePrefix("fd12:3456:7890:abcd::/64"),
|
||||
}
|
||||
|
||||
for _, prefix := range prefixes {
|
||||
t.Run(prefix.String(), func(t *testing.T) {
|
||||
_, err := AllocatePeerIP(prefix, nil)
|
||||
assert.Error(t, err)
|
||||
|
||||
_, err = AllocateRandomPeerIP(prefix)
|
||||
assert.Error(t, err)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAllocatePeerIPIgnoresNonIPv4TakenIPs(t *testing.T) {
|
||||
prefix := netip.MustParsePrefix("192.168.1.0/29")
|
||||
|
||||
ip, err := AllocatePeerIP(prefix, []netip.Addr{netip.MustParseAddr("fd12:3456:7890:abcd::1")})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, prefix.Contains(ip))
|
||||
}
|
||||
|
||||
func TestGenerateIPs(t *testing.T) {
|
||||
ipNet := net.IPNet{IP: net.ParseIP("100.64.0.0"), Mask: net.IPMask{255, 255, 255, 0}}
|
||||
ips, ipsLen := generateIPs(&ipNet, map[string]struct{}{"100.64.0.0": {}})
|
||||
if ipsLen != 252 {
|
||||
t.Errorf("expected 252 ips, got %d", len(ips))
|
||||
return
|
||||
}
|
||||
if ips[len(ips)-1].String() != "100.64.0.253" {
|
||||
t.Errorf("expected last ip to be: 100.64.0.253, got %s", ips[len(ips)-1].String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewNetworkHasIPv6(t *testing.T) {
|
||||
network := NewNetwork()
|
||||
|
||||
assert.NotNil(t, network.NetV6.IP, "v6 subnet should be allocated")
|
||||
assert.True(t, network.NetV6.IP.To4() == nil, "v6 subnet should be IPv6")
|
||||
assert.Equal(t, byte(0xfd), network.NetV6.IP[0], "v6 subnet should be ULA (fd prefix)")
|
||||
|
||||
ones, bits := network.NetV6.Mask.Size()
|
||||
assert.Equal(t, 64, ones, "v6 subnet should be /64")
|
||||
assert.Equal(t, 128, bits)
|
||||
}
|
||||
|
||||
func TestAllocateIPv6SubnetUniqueness(t *testing.T) {
|
||||
seen := make(map[string]struct{})
|
||||
for i := 0; i < 100; i++ {
|
||||
network := NewNetwork()
|
||||
key := network.NetV6.IP.String()
|
||||
_, duplicate := seen[key]
|
||||
assert.False(t, duplicate, "duplicate v6 subnet: %s", key)
|
||||
seen[key] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAllocateRandomPeerIPv6(t *testing.T) {
|
||||
prefix := netip.MustParsePrefix("fd12:3456:7890:abcd::/64")
|
||||
|
||||
ip, err := AllocateRandomPeerIPv6(prefix)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.True(t, ip.Is6(), "should be IPv6")
|
||||
assert.True(t, prefix.Contains(ip), "should be within subnet")
|
||||
// First 8 bytes (network prefix) should match
|
||||
b := ip.As16()
|
||||
prefixBytes := prefix.Addr().As16()
|
||||
assert.Equal(t, prefixBytes[:8], b[:8], "prefix should match")
|
||||
// Interface ID should not be all zeros
|
||||
allZero := true
|
||||
for _, v := range b[8:] {
|
||||
if v != 0 {
|
||||
allZero = false
|
||||
break
|
||||
}
|
||||
}
|
||||
assert.False(t, allZero, "interface ID should not be all zeros")
|
||||
}
|
||||
|
||||
func TestAllocateRandomPeerIPv6_VariousPrefixes(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
cidr string
|
||||
prefix int
|
||||
}{
|
||||
{"standard /64", "fd00:1234:5678:abcd::/64", 64},
|
||||
{"small /112", "fd00:1234:5678:abcd::/112", 112},
|
||||
{"large /48", "fd00:1234::/48", 48},
|
||||
{"non-boundary /60", "fd00:1234:5670::/60", 60},
|
||||
{"non-boundary /52", "fd00:1230::/52", 52},
|
||||
{"minimum /120", "fd00:1234:5678:abcd::100/120", 120},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
prefix, err := netip.ParsePrefix(tt.cidr)
|
||||
require.NoError(t, err)
|
||||
prefix = prefix.Masked()
|
||||
|
||||
assert.Equal(t, tt.prefix, prefix.Bits())
|
||||
|
||||
for i := 0; i < 50; i++ {
|
||||
ip, err := AllocateRandomPeerIPv6(prefix)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, prefix.Contains(ip), "IP %s should be within %s", ip, prefix)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAllocateRandomPeerIPv6_PreservesNetworkBits(t *testing.T) {
|
||||
// For a /112, bytes 0-13 should be preserved, only bytes 14-15 should vary
|
||||
prefix := netip.MustParsePrefix("fd00:1234:5678:abcd:ef01:2345:6789:0/112")
|
||||
|
||||
prefixBytes := prefix.Addr().As16()
|
||||
for i := 0; i < 20; i++ {
|
||||
ip, err := AllocateRandomPeerIPv6(prefix)
|
||||
require.NoError(t, err)
|
||||
// First 14 bytes (112 bits = 14 bytes) must match the network
|
||||
b := ip.As16()
|
||||
assert.Equal(t, prefixBytes[:14], b[:14], "network bytes should be preserved for /112")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAllocateRandomPeerIPv6_NonByteBoundary(t *testing.T) {
|
||||
// For a /60, the first 7.5 bytes are network, so byte 7 is partial
|
||||
prefix := netip.MustParsePrefix("fd00:1234:5678:abc0::/60")
|
||||
|
||||
prefixBytes := prefix.Addr().As16()
|
||||
for i := 0; i < 50; i++ {
|
||||
ip, err := AllocateRandomPeerIPv6(prefix)
|
||||
require.NoError(t, err)
|
||||
b := ip.As16()
|
||||
assert.True(t, prefix.Contains(ip), "IP %s should be within %s", ip, prefix)
|
||||
// First 7 bytes must match exactly
|
||||
assert.Equal(t, prefixBytes[:7], b[:7], "full network bytes should match for /60")
|
||||
// Byte 7: top 4 bits (0xc = 1100) must be preserved
|
||||
assert.Equal(t, prefixBytes[7]&0xf0, b[7]&0xf0, "partial byte network bits should be preserved for /60")
|
||||
}
|
||||
}
|
||||
1032
shared/management/types/networkmap_components.go
Normal file
1032
shared/management/types/networkmap_components.go
Normal file
File diff suppressed because it is too large
Load Diff
227
shared/management/types/networkmap_components_compact.go
Normal file
227
shared/management/types/networkmap_components_compact.go
Normal file
@@ -0,0 +1,227 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
"github.com/netbirdio/netbird/route"
|
||||
)
|
||||
|
||||
type GroupCompact struct {
|
||||
Name string
|
||||
PeerIndexes []int
|
||||
}
|
||||
|
||||
type NetworkMapComponentsCompact struct {
|
||||
PeerID string
|
||||
|
||||
Network *Network
|
||||
AccountSettings *AccountSettingsInfo
|
||||
DNSSettings *DNSSettings
|
||||
CustomZoneDomain string
|
||||
|
||||
AllPeers []*ComponentPeer
|
||||
PeerIndexes []int
|
||||
RouterPeerIndexes []int
|
||||
|
||||
Groups map[string]*GroupCompact
|
||||
AllPolicies []*Policy
|
||||
PolicyIndexes []int
|
||||
ResourcePoliciesMap map[string][]int
|
||||
Routes []*route.Route
|
||||
NameServerGroups []*nbdns.NameServerGroup
|
||||
AllDNSRecords []nbdns.SimpleRecord
|
||||
AccountZones []nbdns.CustomZone
|
||||
|
||||
RoutersMap map[string]map[string]*ComponentRouter
|
||||
NetworkResources []*ComponentResource
|
||||
|
||||
GroupIDToUserIDs map[string][]string
|
||||
AllowedUserIDs map[string]struct{}
|
||||
PostureFailedPeers map[string]map[string]struct{}
|
||||
}
|
||||
|
||||
func (c *NetworkMapComponents) ToCompact() *NetworkMapComponentsCompact {
|
||||
peerToIndex := make(map[string]int)
|
||||
var allPeers []*ComponentPeer
|
||||
|
||||
for id, peer := range c.Peers {
|
||||
if _, exists := peerToIndex[id]; !exists {
|
||||
peerToIndex[id] = len(allPeers)
|
||||
allPeers = append(allPeers, peer)
|
||||
}
|
||||
}
|
||||
|
||||
for id, peer := range c.RouterPeers {
|
||||
if _, exists := peerToIndex[id]; !exists {
|
||||
peerToIndex[id] = len(allPeers)
|
||||
allPeers = append(allPeers, peer)
|
||||
}
|
||||
}
|
||||
|
||||
peerIndexes := make([]int, 0, len(c.Peers))
|
||||
for id := range c.Peers {
|
||||
peerIndexes = append(peerIndexes, peerToIndex[id])
|
||||
}
|
||||
|
||||
routerPeerIndexes := make([]int, 0, len(c.RouterPeers))
|
||||
for id := range c.RouterPeers {
|
||||
routerPeerIndexes = append(routerPeerIndexes, peerToIndex[id])
|
||||
}
|
||||
|
||||
groups := make(map[string]*GroupCompact, len(c.Groups))
|
||||
for id, group := range c.Groups {
|
||||
peerIdxs := make([]int, 0, len(group.Peers))
|
||||
for _, peerID := range group.Peers {
|
||||
if idx, ok := peerToIndex[peerID]; ok {
|
||||
peerIdxs = append(peerIdxs, idx)
|
||||
}
|
||||
}
|
||||
groups[id] = &GroupCompact{
|
||||
Name: group.Name,
|
||||
PeerIndexes: peerIdxs,
|
||||
}
|
||||
}
|
||||
|
||||
policyToIndex := make(map[*Policy]int)
|
||||
var allPolicies []*Policy
|
||||
|
||||
for _, policy := range c.Policies {
|
||||
if _, exists := policyToIndex[policy]; !exists {
|
||||
policyToIndex[policy] = len(allPolicies)
|
||||
allPolicies = append(allPolicies, policy)
|
||||
}
|
||||
}
|
||||
|
||||
for _, policies := range c.ResourcePoliciesMap {
|
||||
for _, policy := range policies {
|
||||
if _, exists := policyToIndex[policy]; !exists {
|
||||
policyToIndex[policy] = len(allPolicies)
|
||||
allPolicies = append(allPolicies, policy)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
policyIndexes := make([]int, len(c.Policies))
|
||||
for i, policy := range c.Policies {
|
||||
policyIndexes[i] = policyToIndex[policy]
|
||||
}
|
||||
|
||||
var resourcePoliciesMap map[string][]int
|
||||
if len(c.ResourcePoliciesMap) > 0 {
|
||||
resourcePoliciesMap = make(map[string][]int, len(c.ResourcePoliciesMap))
|
||||
for resID, policies := range c.ResourcePoliciesMap {
|
||||
indexes := make([]int, len(policies))
|
||||
for i, policy := range policies {
|
||||
indexes[i] = policyToIndex[policy]
|
||||
}
|
||||
resourcePoliciesMap[resID] = indexes
|
||||
}
|
||||
}
|
||||
|
||||
return &NetworkMapComponentsCompact{
|
||||
PeerID: c.PeerID,
|
||||
Network: c.Network,
|
||||
AccountSettings: c.AccountSettings,
|
||||
DNSSettings: c.DNSSettings,
|
||||
CustomZoneDomain: c.CustomZoneDomain,
|
||||
|
||||
AllPeers: allPeers,
|
||||
PeerIndexes: peerIndexes,
|
||||
RouterPeerIndexes: routerPeerIndexes,
|
||||
|
||||
Groups: groups,
|
||||
AllPolicies: allPolicies,
|
||||
PolicyIndexes: policyIndexes,
|
||||
ResourcePoliciesMap: resourcePoliciesMap,
|
||||
Routes: c.Routes,
|
||||
NameServerGroups: c.NameServerGroups,
|
||||
AllDNSRecords: c.AllDNSRecords,
|
||||
AccountZones: c.AccountZones,
|
||||
|
||||
RoutersMap: c.RoutersMap,
|
||||
NetworkResources: c.NetworkResources,
|
||||
|
||||
GroupIDToUserIDs: c.GroupIDToUserIDs,
|
||||
AllowedUserIDs: c.AllowedUserIDs,
|
||||
PostureFailedPeers: c.PostureFailedPeers,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *NetworkMapComponentsCompact) ToFull() *NetworkMapComponents {
|
||||
peers := make(map[string]*ComponentPeer, len(c.PeerIndexes))
|
||||
for _, idx := range c.PeerIndexes {
|
||||
if idx >= 0 && idx < len(c.AllPeers) {
|
||||
peer := c.AllPeers[idx]
|
||||
peers[peer.ID] = peer
|
||||
}
|
||||
}
|
||||
|
||||
routerPeers := make(map[string]*ComponentPeer, len(c.RouterPeerIndexes))
|
||||
for _, idx := range c.RouterPeerIndexes {
|
||||
if idx >= 0 && idx < len(c.AllPeers) {
|
||||
peer := c.AllPeers[idx]
|
||||
routerPeers[peer.ID] = peer
|
||||
}
|
||||
}
|
||||
|
||||
groups := make(map[string]*ComponentGroup, len(c.Groups))
|
||||
for id, gc := range c.Groups {
|
||||
peerIDs := make([]string, 0, len(gc.PeerIndexes))
|
||||
for _, idx := range gc.PeerIndexes {
|
||||
if idx >= 0 && idx < len(c.AllPeers) {
|
||||
peerIDs = append(peerIDs, c.AllPeers[idx].ID)
|
||||
}
|
||||
}
|
||||
groups[id] = &ComponentGroup{
|
||||
ID: id,
|
||||
Name: gc.Name,
|
||||
Peers: peerIDs,
|
||||
}
|
||||
}
|
||||
|
||||
policies := make([]*Policy, len(c.PolicyIndexes))
|
||||
for i, idx := range c.PolicyIndexes {
|
||||
if idx >= 0 && idx < len(c.AllPolicies) {
|
||||
policies[i] = c.AllPolicies[idx]
|
||||
}
|
||||
}
|
||||
|
||||
var resourcePoliciesMap map[string][]*Policy
|
||||
if len(c.ResourcePoliciesMap) > 0 {
|
||||
resourcePoliciesMap = make(map[string][]*Policy, len(c.ResourcePoliciesMap))
|
||||
for resID, indexes := range c.ResourcePoliciesMap {
|
||||
pols := make([]*Policy, 0, len(indexes))
|
||||
for _, idx := range indexes {
|
||||
if idx >= 0 && idx < len(c.AllPolicies) {
|
||||
pols = append(pols, c.AllPolicies[idx])
|
||||
}
|
||||
}
|
||||
resourcePoliciesMap[resID] = pols
|
||||
}
|
||||
}
|
||||
|
||||
return &NetworkMapComponents{
|
||||
PeerID: c.PeerID,
|
||||
Network: c.Network,
|
||||
AccountSettings: c.AccountSettings,
|
||||
DNSSettings: c.DNSSettings,
|
||||
CustomZoneDomain: c.CustomZoneDomain,
|
||||
|
||||
Peers: peers,
|
||||
RouterPeers: routerPeers,
|
||||
|
||||
Groups: groups,
|
||||
Policies: policies,
|
||||
Routes: c.Routes,
|
||||
NameServerGroups: c.NameServerGroups,
|
||||
AllDNSRecords: c.AllDNSRecords,
|
||||
AccountZones: c.AccountZones,
|
||||
|
||||
ResourcePoliciesMap: resourcePoliciesMap,
|
||||
RoutersMap: c.RoutersMap,
|
||||
NetworkResources: c.NetworkResources,
|
||||
|
||||
GroupIDToUserIDs: c.GroupIDToUserIDs,
|
||||
AllowedUserIDs: c.AllowedUserIDs,
|
||||
PostureFailedPeers: c.PostureFailedPeers,
|
||||
}
|
||||
}
|
||||
268
shared/management/types/policy.go
Normal file
268
shared/management/types/policy.go
Normal file
@@ -0,0 +1,268 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
// PolicyTrafficActionAccept indicates that the traffic is accepted
|
||||
PolicyTrafficActionAccept = PolicyTrafficActionType("accept")
|
||||
// PolicyTrafficActionDrop indicates that the traffic is dropped
|
||||
PolicyTrafficActionDrop = PolicyTrafficActionType("drop")
|
||||
)
|
||||
|
||||
const (
|
||||
// PolicyRuleProtocolALL type of traffic
|
||||
PolicyRuleProtocolALL = PolicyRuleProtocolType("all")
|
||||
// PolicyRuleProtocolTCP type of traffic
|
||||
PolicyRuleProtocolTCP = PolicyRuleProtocolType("tcp")
|
||||
// PolicyRuleProtocolUDP type of traffic
|
||||
PolicyRuleProtocolUDP = PolicyRuleProtocolType("udp")
|
||||
// PolicyRuleProtocolICMP type of traffic
|
||||
PolicyRuleProtocolICMP = PolicyRuleProtocolType("icmp")
|
||||
// PolicyRuleProtocolNetbirdSSH type of traffic
|
||||
PolicyRuleProtocolNetbirdSSH = PolicyRuleProtocolType("netbird-ssh")
|
||||
)
|
||||
|
||||
const (
|
||||
// PolicyRuleFlowDirect allows traffic from source to destination
|
||||
PolicyRuleFlowDirect = PolicyRuleDirection("direct")
|
||||
// PolicyRuleFlowBidirect allows traffic to both directions
|
||||
PolicyRuleFlowBidirect = PolicyRuleDirection("bidirect")
|
||||
)
|
||||
|
||||
const (
|
||||
// DefaultRuleName is a name for the Default rule that is created for every account
|
||||
DefaultRuleName = "Default"
|
||||
// DefaultRuleDescription is a description for the Default rule that is created for every account
|
||||
DefaultRuleDescription = "This is a default rule that allows connections between all the resources"
|
||||
// DefaultPolicyName is a name for the Default policy that is created for every account
|
||||
DefaultPolicyName = "Default"
|
||||
// DefaultPolicyDescription is a description for the Default policy that is created for every account
|
||||
DefaultPolicyDescription = "This is a default policy that allows connections between all the resources"
|
||||
)
|
||||
|
||||
// PolicyUpdateOperation operation object with type and values to be applied
|
||||
type PolicyUpdateOperation struct {
|
||||
Type PolicyUpdateOperationType
|
||||
Values []string
|
||||
}
|
||||
|
||||
// Policy of the Rego query
|
||||
type Policy struct {
|
||||
// ID of the policy'
|
||||
ID string `gorm:"primaryKey"`
|
||||
|
||||
PublicID string `json:"-"`
|
||||
|
||||
// AccountID is a reference to Account that this object belongs
|
||||
AccountID string `json:"-" gorm:"index"`
|
||||
|
||||
// Name of the Policy
|
||||
Name string
|
||||
|
||||
// Description of the policy visible in the UI
|
||||
Description string
|
||||
|
||||
// Enabled status of the policy
|
||||
Enabled bool
|
||||
|
||||
// Rules of the policy
|
||||
Rules []*PolicyRule `gorm:"foreignKey:PolicyID;references:id;constraint:OnDelete:CASCADE;"`
|
||||
|
||||
// SourcePostureChecks are ID references to Posture checks for policy source groups
|
||||
SourcePostureChecks []string `gorm:"serializer:json"`
|
||||
}
|
||||
|
||||
// Copy returns a copy of the policy.
|
||||
func (p *Policy) Copy() *Policy {
|
||||
c := &Policy{
|
||||
ID: p.ID,
|
||||
AccountID: p.AccountID,
|
||||
PublicID: p.PublicID,
|
||||
Name: p.Name,
|
||||
Description: p.Description,
|
||||
Enabled: p.Enabled,
|
||||
Rules: make([]*PolicyRule, len(p.Rules)),
|
||||
SourcePostureChecks: make([]string, len(p.SourcePostureChecks)),
|
||||
}
|
||||
for i, r := range p.Rules {
|
||||
c.Rules[i] = r.Copy()
|
||||
}
|
||||
copy(c.SourcePostureChecks, p.SourcePostureChecks)
|
||||
return c
|
||||
}
|
||||
|
||||
func (p *Policy) Equal(other *Policy) bool {
|
||||
if p == nil || other == nil {
|
||||
return p == other
|
||||
}
|
||||
|
||||
if p.ID != other.ID ||
|
||||
p.AccountID != other.AccountID ||
|
||||
p.Name != other.Name ||
|
||||
p.Description != other.Description ||
|
||||
p.Enabled != other.Enabled {
|
||||
return false
|
||||
}
|
||||
|
||||
if !stringSlicesEqualUnordered(p.SourcePostureChecks, other.SourcePostureChecks) {
|
||||
return false
|
||||
}
|
||||
|
||||
if len(p.Rules) != len(other.Rules) {
|
||||
return false
|
||||
}
|
||||
|
||||
otherRules := make(map[string]*PolicyRule, len(other.Rules))
|
||||
for _, r := range other.Rules {
|
||||
otherRules[r.ID] = r
|
||||
}
|
||||
for _, r := range p.Rules {
|
||||
otherRule, ok := otherRules[r.ID]
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
if !r.Equal(otherRule) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
// EventMeta returns activity event meta related to this policy
|
||||
func (p *Policy) EventMeta() map[string]any {
|
||||
return map[string]any{"name": p.Name}
|
||||
}
|
||||
|
||||
// UpgradeAndFix different version of policies to latest version
|
||||
func (p *Policy) UpgradeAndFix() {
|
||||
for _, r := range p.Rules {
|
||||
// start migrate from version v0.20.3
|
||||
if r.Protocol == "" {
|
||||
r.Protocol = PolicyRuleProtocolALL
|
||||
}
|
||||
if r.Protocol == PolicyRuleProtocolALL && !r.Bidirectional {
|
||||
r.Bidirectional = true
|
||||
}
|
||||
// -- v0.20.4
|
||||
}
|
||||
}
|
||||
|
||||
// RuleGroups returns a list of all groups referenced in the policy's rules,
|
||||
// including sources and destinations.
|
||||
func (p *Policy) RuleGroups() []string {
|
||||
groups := make([]string, 0)
|
||||
for _, rule := range p.Rules {
|
||||
groups = append(groups, rule.Sources...)
|
||||
groups = append(groups, rule.Destinations...)
|
||||
}
|
||||
|
||||
return groups
|
||||
}
|
||||
|
||||
// SourceGroups returns a slice of all unique source groups referenced in the policy's rules.
|
||||
func (p *Policy) SourceGroups() []string {
|
||||
if len(p.Rules) == 1 {
|
||||
return p.Rules[0].Sources
|
||||
}
|
||||
groups := make(map[string]struct{}, len(p.Rules))
|
||||
for _, rule := range p.Rules {
|
||||
for _, source := range rule.Sources {
|
||||
groups[source] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
groupIDs := make([]string, 0, len(groups))
|
||||
for groupID := range groups {
|
||||
groupIDs = append(groupIDs, groupID)
|
||||
}
|
||||
|
||||
return groupIDs
|
||||
}
|
||||
|
||||
func ParseRuleString(rule string) (PolicyRuleProtocolType, RulePortRange, error) {
|
||||
rule = strings.TrimSpace(strings.ToLower(rule))
|
||||
if rule == "all" {
|
||||
return PolicyRuleProtocolALL, RulePortRange{}, nil
|
||||
}
|
||||
if rule == "icmp" {
|
||||
return PolicyRuleProtocolICMP, RulePortRange{}, nil
|
||||
}
|
||||
|
||||
split := strings.Split(rule, "/")
|
||||
if len(split) != 2 {
|
||||
return "", RulePortRange{}, errors.New("invalid rule format: expected protocol/port or protocol/port-range")
|
||||
}
|
||||
|
||||
protoStr := strings.TrimSpace(split[0])
|
||||
portStr := strings.TrimSpace(split[1])
|
||||
|
||||
var protocol PolicyRuleProtocolType
|
||||
switch protoStr {
|
||||
case "tcp":
|
||||
protocol = PolicyRuleProtocolTCP
|
||||
case "udp":
|
||||
protocol = PolicyRuleProtocolUDP
|
||||
case "icmp":
|
||||
return "", RulePortRange{}, errors.New("icmp does not accept ports; use 'icmp' without '/…'")
|
||||
case "netbird-ssh":
|
||||
return PolicyRuleProtocolNetbirdSSH, RulePortRange{Start: nativeSSHPortNumber, End: nativeSSHPortNumber}, nil
|
||||
default:
|
||||
return "", RulePortRange{}, fmt.Errorf("invalid protocol: %q", protoStr)
|
||||
}
|
||||
|
||||
portRange, err := parsePortRange(portStr)
|
||||
if err != nil {
|
||||
return "", RulePortRange{}, err
|
||||
}
|
||||
|
||||
return protocol, portRange, nil
|
||||
}
|
||||
|
||||
func parsePortRange(portStr string) (RulePortRange, error) {
|
||||
if strings.Contains(portStr, "-") {
|
||||
rangeParts := strings.Split(portStr, "-")
|
||||
if len(rangeParts) != 2 {
|
||||
return RulePortRange{}, fmt.Errorf("invalid port range %q", portStr)
|
||||
}
|
||||
start, err := parsePort(strings.TrimSpace(rangeParts[0]))
|
||||
if err != nil {
|
||||
return RulePortRange{}, err
|
||||
}
|
||||
end, err := parsePort(strings.TrimSpace(rangeParts[1]))
|
||||
if err != nil {
|
||||
return RulePortRange{}, err
|
||||
}
|
||||
if start > end {
|
||||
return RulePortRange{}, fmt.Errorf("invalid port range: start %d > end %d", start, end)
|
||||
}
|
||||
return RulePortRange{Start: uint16(start), End: uint16(end)}, nil
|
||||
}
|
||||
|
||||
p, err := parsePort(portStr)
|
||||
if err != nil {
|
||||
return RulePortRange{}, err
|
||||
}
|
||||
|
||||
return RulePortRange{Start: uint16(p), End: uint16(p)}, nil
|
||||
}
|
||||
|
||||
func parsePort(portStr string) (int, error) {
|
||||
|
||||
if portStr == "" {
|
||||
return 0, errors.New("empty port")
|
||||
}
|
||||
p, err := strconv.Atoi(portStr)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("invalid port %q: %w", portStr, err)
|
||||
}
|
||||
if p < 1 || p > 65535 {
|
||||
return 0, fmt.Errorf("port out of range (1–65535): %d", p)
|
||||
}
|
||||
return p, nil
|
||||
}
|
||||
225
shared/management/types/policyrule.go
Normal file
225
shared/management/types/policyrule.go
Normal file
@@ -0,0 +1,225 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"slices"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
// PolicyUpdateOperationType operation type
|
||||
type PolicyUpdateOperationType int
|
||||
|
||||
// PolicyTrafficActionType action type for the firewall
|
||||
type PolicyTrafficActionType string
|
||||
|
||||
// PolicyRuleProtocolType type of traffic
|
||||
type PolicyRuleProtocolType string
|
||||
|
||||
// PolicyRuleDirection direction of traffic
|
||||
type PolicyRuleDirection string
|
||||
|
||||
// RulePortRange represents a range of ports for a firewall rule.
|
||||
type RulePortRange struct {
|
||||
Start uint16
|
||||
End uint16
|
||||
}
|
||||
|
||||
func (r *RulePortRange) ToProto() *proto.PortInfo {
|
||||
return &proto.PortInfo{
|
||||
PortSelection: &proto.PortInfo_Range_{
|
||||
Range: &proto.PortInfo_Range{
|
||||
Start: uint32(r.Start),
|
||||
End: uint32(r.End),
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (r *RulePortRange) Equal(other *RulePortRange) bool {
|
||||
return r.Start == other.Start && r.End == other.End
|
||||
}
|
||||
|
||||
// PolicyRule is the metadata of the policy
|
||||
type PolicyRule struct {
|
||||
// ID of the policy rule
|
||||
ID string `gorm:"primaryKey"`
|
||||
|
||||
// PolicyID is a reference to Policy that this object belongs
|
||||
PolicyID string `json:"-" gorm:"index"`
|
||||
|
||||
// Name of the rule visible in the UI
|
||||
Name string
|
||||
|
||||
// Description of the rule visible in the UI
|
||||
Description string
|
||||
|
||||
// Enabled status of rule in the system
|
||||
Enabled bool
|
||||
|
||||
// Action policy accept or drops packets
|
||||
Action PolicyTrafficActionType
|
||||
|
||||
// Destinations policy destination groups
|
||||
Destinations []string `gorm:"serializer:json"`
|
||||
|
||||
// DestinationResource policy destination resource that the rule is applied to
|
||||
DestinationResource Resource `gorm:"serializer:json"`
|
||||
|
||||
// Sources policy source groups
|
||||
Sources []string `gorm:"serializer:json"`
|
||||
|
||||
// SourceResource policy source resource that the rule is applied to
|
||||
SourceResource Resource `gorm:"serializer:json"`
|
||||
|
||||
// Bidirectional define if the rule is applicable in both directions, sources, and destinations
|
||||
Bidirectional bool
|
||||
|
||||
// Protocol type of the traffic
|
||||
Protocol PolicyRuleProtocolType
|
||||
|
||||
// Ports or it ranges list
|
||||
Ports []string `gorm:"serializer:json"`
|
||||
|
||||
// PortRanges a list of port ranges.
|
||||
PortRanges []RulePortRange `gorm:"serializer:json"`
|
||||
|
||||
// AuthorizedGroups is a map of groupIDs and their respective access to local users via ssh
|
||||
AuthorizedGroups map[string][]string `gorm:"serializer:json"`
|
||||
|
||||
// AuthorizedUser is a list of userIDs that are authorized to access local resources via ssh
|
||||
AuthorizedUser string
|
||||
}
|
||||
|
||||
// Copy returns a copy of a policy rule
|
||||
func (pm *PolicyRule) Copy() *PolicyRule {
|
||||
rule := &PolicyRule{
|
||||
ID: pm.ID,
|
||||
PolicyID: pm.PolicyID,
|
||||
Name: pm.Name,
|
||||
Description: pm.Description,
|
||||
Enabled: pm.Enabled,
|
||||
Action: pm.Action,
|
||||
Destinations: make([]string, len(pm.Destinations)),
|
||||
DestinationResource: pm.DestinationResource,
|
||||
Sources: make([]string, len(pm.Sources)),
|
||||
SourceResource: pm.SourceResource,
|
||||
Bidirectional: pm.Bidirectional,
|
||||
Protocol: pm.Protocol,
|
||||
Ports: make([]string, len(pm.Ports)),
|
||||
PortRanges: make([]RulePortRange, len(pm.PortRanges)),
|
||||
AuthorizedGroups: make(map[string][]string, len(pm.AuthorizedGroups)),
|
||||
AuthorizedUser: pm.AuthorizedUser,
|
||||
}
|
||||
copy(rule.Destinations, pm.Destinations)
|
||||
copy(rule.Sources, pm.Sources)
|
||||
copy(rule.Ports, pm.Ports)
|
||||
copy(rule.PortRanges, pm.PortRanges)
|
||||
for k, v := range pm.AuthorizedGroups {
|
||||
rule.AuthorizedGroups[k] = make([]string, len(v))
|
||||
copy(rule.AuthorizedGroups[k], v)
|
||||
}
|
||||
return rule
|
||||
}
|
||||
|
||||
func (pm *PolicyRule) Equal(other *PolicyRule) bool {
|
||||
if pm == nil || other == nil {
|
||||
return pm == other
|
||||
}
|
||||
|
||||
if pm.ID != other.ID ||
|
||||
pm.PolicyID != other.PolicyID ||
|
||||
pm.Name != other.Name ||
|
||||
pm.Description != other.Description ||
|
||||
pm.Enabled != other.Enabled ||
|
||||
pm.Action != other.Action ||
|
||||
pm.Bidirectional != other.Bidirectional ||
|
||||
pm.Protocol != other.Protocol ||
|
||||
pm.SourceResource != other.SourceResource ||
|
||||
pm.DestinationResource != other.DestinationResource ||
|
||||
pm.AuthorizedUser != other.AuthorizedUser {
|
||||
return false
|
||||
}
|
||||
|
||||
if !stringSlicesEqualUnordered(pm.Sources, other.Sources) {
|
||||
return false
|
||||
}
|
||||
if !stringSlicesEqualUnordered(pm.Destinations, other.Destinations) {
|
||||
return false
|
||||
}
|
||||
if !stringSlicesEqualUnordered(pm.Ports, other.Ports) {
|
||||
return false
|
||||
}
|
||||
if !portRangeSlicesEqualUnordered(pm.PortRanges, other.PortRanges) {
|
||||
return false
|
||||
}
|
||||
if !authorizedGroupsEqual(pm.AuthorizedGroups, other.AuthorizedGroups) {
|
||||
return false
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
func stringSlicesEqualUnordered(a, b []string) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
if len(a) == 0 {
|
||||
return true
|
||||
}
|
||||
sorted1 := make([]string, len(a))
|
||||
sorted2 := make([]string, len(b))
|
||||
copy(sorted1, a)
|
||||
copy(sorted2, b)
|
||||
slices.Sort(sorted1)
|
||||
slices.Sort(sorted2)
|
||||
return slices.Equal(sorted1, sorted2)
|
||||
}
|
||||
|
||||
func portRangeSlicesEqualUnordered(a, b []RulePortRange) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
if len(a) == 0 {
|
||||
return true
|
||||
}
|
||||
cmp := func(x, y RulePortRange) int {
|
||||
if x.Start != y.Start {
|
||||
if x.Start < y.Start {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
}
|
||||
if x.End != y.End {
|
||||
if x.End < y.End {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
sorted1 := make([]RulePortRange, len(a))
|
||||
sorted2 := make([]RulePortRange, len(b))
|
||||
copy(sorted1, a)
|
||||
copy(sorted2, b)
|
||||
slices.SortFunc(sorted1, cmp)
|
||||
slices.SortFunc(sorted2, cmp)
|
||||
return slices.EqualFunc(sorted1, sorted2, func(x, y RulePortRange) bool {
|
||||
return x.Start == y.Start && x.End == y.End
|
||||
})
|
||||
}
|
||||
|
||||
func authorizedGroupsEqual(a, b map[string][]string) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
for k, va := range a {
|
||||
vb, ok := b[k]
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
if !stringSlicesEqualUnordered(va, vb) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
39
shared/management/types/resource.go
Normal file
39
shared/management/types/resource.go
Normal file
@@ -0,0 +1,39 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
)
|
||||
|
||||
type ResourceType string
|
||||
|
||||
const (
|
||||
ResourceTypePeer ResourceType = "peer"
|
||||
ResourceTypeDomain ResourceType = "domain"
|
||||
ResourceTypeHost ResourceType = "host"
|
||||
ResourceTypeSubnet ResourceType = "subnet"
|
||||
)
|
||||
|
||||
type Resource struct {
|
||||
ID string
|
||||
Type ResourceType
|
||||
}
|
||||
|
||||
func (r *Resource) ToAPIResponse() *api.Resource {
|
||||
if r.ID == "" && r.Type == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
return &api.Resource{
|
||||
Id: r.ID,
|
||||
Type: api.ResourceType(r.Type),
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Resource) FromAPIRequest(req *api.Resource) {
|
||||
if req == nil {
|
||||
return
|
||||
}
|
||||
|
||||
r.ID = req.Id
|
||||
r.Type = ResourceType(req.Type)
|
||||
}
|
||||
64
shared/management/types/route_firewall_rule.go
Normal file
64
shared/management/types/route_firewall_rule.go
Normal file
@@ -0,0 +1,64 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"github.com/netbirdio/netbird/route"
|
||||
"github.com/netbirdio/netbird/shared/management/domain"
|
||||
)
|
||||
|
||||
// RouteFirewallRule a firewall rule applicable for a routed network.
|
||||
type RouteFirewallRule struct {
|
||||
// PolicyID is the ID of the policy this rule is derived from
|
||||
PolicyID string
|
||||
|
||||
// RouteID is the ID of the route this rule belongs to.
|
||||
RouteID route.ID
|
||||
|
||||
// SourceRanges IP ranges of the routing peers.
|
||||
SourceRanges []string
|
||||
|
||||
// Action of the traffic when the rule is applicable
|
||||
Action string
|
||||
|
||||
// Destination a network prefix for the routed traffic
|
||||
Destination string
|
||||
|
||||
// Protocol of the traffic
|
||||
Protocol string
|
||||
|
||||
// Port of the traffic
|
||||
Port uint16
|
||||
|
||||
// PortRange represents the range of ports for a firewall rule
|
||||
PortRange RulePortRange
|
||||
|
||||
// Domains list of network domains for the routed traffic
|
||||
Domains domain.List
|
||||
|
||||
// isDynamic indicates whether the rule is for DNS routing
|
||||
IsDynamic bool
|
||||
}
|
||||
|
||||
func (r *RouteFirewallRule) Equal(other *RouteFirewallRule) bool {
|
||||
if r.Action != other.Action {
|
||||
return false
|
||||
}
|
||||
if r.Destination != other.Destination {
|
||||
return false
|
||||
}
|
||||
if r.Protocol != other.Protocol {
|
||||
return false
|
||||
}
|
||||
if r.Port != other.Port {
|
||||
return false
|
||||
}
|
||||
if !r.PortRange.Equal(&other.PortRange) {
|
||||
return false
|
||||
}
|
||||
if !r.Domains.Equal(other.Domains) {
|
||||
return false
|
||||
}
|
||||
if r.IsDynamic != other.IsDynamic {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
Reference in New Issue
Block a user