mirror of
https://github.com/netbirdio/netbird.git
synced 2026-08-03 22:31:30 +02:00
Compare commits
36 Commits
mdm_integr
...
revert/com
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2bb55b6c88 | ||
|
|
b8e004ea89 | ||
|
|
e2797360f4 | ||
|
|
2f399f1e6e | ||
|
|
82c1e18264 | ||
|
|
ce9023bd27 | ||
|
|
7d33356776 | ||
|
|
993291149c | ||
|
|
98f7ea40a1 | ||
|
|
76db9ab94f | ||
|
|
9653b15d78 | ||
|
|
206bb1676b | ||
|
|
600b0c752b | ||
|
|
b02736adc3 | ||
|
|
5d6117d2c0 | ||
|
|
c2f8360b00 | ||
|
|
ea1b4d56e8 | ||
|
|
acbe22b831 | ||
|
|
3b5c8e2298 | ||
|
|
2baeb4bc0d | ||
|
|
4bf75fdc97 | ||
|
|
799f3a3c62 | ||
|
|
129736ad61 | ||
|
|
f7be9c4347 | ||
|
|
e05cb5264d | ||
|
|
5a10561ca1 | ||
|
|
42ce83a8f3 | ||
|
|
e620c86cd4 | ||
|
|
9dee2d60b9 | ||
|
|
4525014632 | ||
|
|
23a5c0de4b | ||
|
|
25e882004f | ||
|
|
15003258d2 | ||
|
|
2af3a5fba5 | ||
|
|
a6603a2e0a | ||
|
|
7ed3737cda |
2
go.mod
2
go.mod
@@ -80,7 +80,7 @@ require (
|
||||
github.com/miekg/dns v1.1.72
|
||||
github.com/mitchellh/hashstructure/v2 v2.0.2
|
||||
github.com/moby/moby/api v1.54.1
|
||||
github.com/netbirdio/management-integrations/integrations v0.0.0-20260416123949-2355d972be42
|
||||
github.com/netbirdio/management-integrations/integrations v0.0.0-20260803100840-78e79ba20f87
|
||||
github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45
|
||||
github.com/oapi-codegen/runtime v1.1.2
|
||||
github.com/okta/okta-sdk-golang/v2 v2.18.0
|
||||
|
||||
2
go.sum
2
go.sum
@@ -486,6 +486,8 @@ github.com/netbirdio/ice/v4 v4.0.0-20250908184934-6202be846b51 h1:Ov4qdafATOgGMB
|
||||
github.com/netbirdio/ice/v4 v4.0.0-20250908184934-6202be846b51/go.mod h1:ZSIbPdBn5hePO8CpF1PekH2SfpTxg1PDhEwtbqZS7R8=
|
||||
github.com/netbirdio/management-integrations/integrations v0.0.0-20260416123949-2355d972be42 h1:F3zS5fT9xzD1OFLfcdAE+3FfyiwjGukF1hvj0jErgs8=
|
||||
github.com/netbirdio/management-integrations/integrations v0.0.0-20260416123949-2355d972be42/go.mod h1:n47r67ZSPgwSmT/Z1o48JjZQW9YJ6m/6Bd/uAXkL3Pg=
|
||||
github.com/netbirdio/management-integrations/integrations v0.0.0-20260803100840-78e79ba20f87 h1:iJeUvSMC0BTpkw7u4JyWcY4/3dl7fEL9DR/TpKf2+1w=
|
||||
github.com/netbirdio/management-integrations/integrations v0.0.0-20260803100840-78e79ba20f87/go.mod h1:pmsCPx1S0nuZRxCextGpc9AV4hLgGSuTsc4NMuwGeCo=
|
||||
github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502 h1:3tHlFmhTdX9axERMVN63dqyFqnvuD+EMJHzM7mNGON8=
|
||||
github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502/go.mod h1:CIMRFEJVL+0DS1a3Nx06NaMn4Dz63Ng6O7dl0qH0zVM=
|
||||
github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45 h1:ujgviVYmx243Ksy7NdSwrdGPSRNE3pb8kEDSpH0QuAQ=
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestGetGroups(t *testing.T) {
|
||||
ctx := context.TODO()
|
||||
|
||||
s, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn)
|
||||
assert.NoError(t, err)
|
||||
|
||||
_, err = s.Pool.Query(ctx,
|
||||
"insert into accounts (id) VALUES('account-id-1')")
|
||||
assert.NoError(t, err)
|
||||
|
||||
_, err = s.Pool.Query(ctx,
|
||||
"insert into groups (id, account_id, name, resources, public_id) VALUES('test-group-id-1','account-id-1','test-group-1', '[{\"ID\":\"host-id-1\",\"Type\":\"host\"}]','public-id-1')")
|
||||
assert.NoError(t, err)
|
||||
_, err = s.Pool.Query(ctx,
|
||||
"insert into groups (id, account_id, name, resources, public_id) VALUES('test-group-id-2','account-id-1','test-group-2', '[{\"ID\":\"subnet-id-1\",\"Type\":\"subnet\"}, {\"ID\":\"host-id-2\",\"Type\":\"host\"}]','public-id-2')")
|
||||
assert.NoError(t, err)
|
||||
|
||||
groups, err := s.GetGroups(ctx, "account-id-1")
|
||||
assert.NoError(t, err)
|
||||
assert.Contains(t,
|
||||
groups,
|
||||
nmdata.Group{Name: "test-group-1", PublicID: "public-id-1", Resources: []nmdata.Resource{{ID: "host-id-1", Type: "host"}}},
|
||||
)
|
||||
assert.Contains(t,
|
||||
groups,
|
||||
nmdata.Group{Name: "test-group-2", PublicID: "public-id-2", Resources: []nmdata.Resource{{ID: "subnet-id-1", Type: "subnet"}, {ID: "host-id-2", Type: "host"}}},
|
||||
)
|
||||
}
|
||||
115
integration_tests/management/network_map_db/pgsql/main_test.go
Normal file
115
integration_tests/management/network_map_db/pgsql/main_test.go
Normal file
@@ -0,0 +1,115 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"regexp"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"gorm.io/driver/postgres"
|
||||
"gorm.io/gorm"
|
||||
|
||||
gormstore "github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/management/server/testutil"
|
||||
)
|
||||
|
||||
var dsn string
|
||||
|
||||
func TestMain(m *testing.M) {
|
||||
_, tmpdsn, err := testutil.CreatePostgresTestContainer()
|
||||
if err != nil {
|
||||
log.Fatalf("error starting postres container %v", err)
|
||||
}
|
||||
|
||||
var db *gorm.DB
|
||||
for i := range 5 {
|
||||
db, err = gorm.Open(postgres.Open(tmpdsn), &gorm.Config{})
|
||||
|
||||
if err == nil {
|
||||
break
|
||||
}
|
||||
|
||||
if i < 5 {
|
||||
waitTime := time.Duration(100*(i+1)) * time.Millisecond
|
||||
time.Sleep(waitTime)
|
||||
continue
|
||||
}
|
||||
|
||||
log.Fatalf("error connecting to postres db %v", err)
|
||||
}
|
||||
|
||||
var cleanup func()
|
||||
dsn, cleanup, err = createRandomDB(tmpdsn, db)
|
||||
sqlDB, _ := db.DB()
|
||||
if sqlDB != nil {
|
||||
sqlDB.Close()
|
||||
}
|
||||
if err != nil {
|
||||
log.Fatalf("error creating postres db %v", err)
|
||||
}
|
||||
|
||||
_, err = gormstore.NewPostgresqlStoreForTests(context.TODO(), dsn, nil, false)
|
||||
if err != nil {
|
||||
log.Fatalf("error running migrations %v", err)
|
||||
}
|
||||
|
||||
code := m.Run()
|
||||
|
||||
cleanup()
|
||||
|
||||
os.Exit(code)
|
||||
}
|
||||
|
||||
func createRandomDB(dsn string, db *gorm.DB) (string, func(), error) {
|
||||
dbName := fmt.Sprintf("test_db_%s", strings.ReplaceAll(uuid.New().String(), "-", "_"))
|
||||
|
||||
if err := db.Exec(fmt.Sprintf("CREATE DATABASE %s", dbName)).Error; err != nil {
|
||||
return "", nil, fmt.Errorf("failed to create database: %v", err)
|
||||
}
|
||||
|
||||
originalDSN := dsn
|
||||
|
||||
cleanup := func() {
|
||||
var dropDB *gorm.DB
|
||||
var err error
|
||||
|
||||
dropDB, err = gorm.Open(postgres.Open(originalDSN), &gorm.Config{
|
||||
SkipDefaultTransaction: true,
|
||||
PrepareStmt: false,
|
||||
})
|
||||
if err != nil {
|
||||
log.Errorf("failed to connect for dropping database %s: %v", dbName, err)
|
||||
return
|
||||
}
|
||||
defer func() {
|
||||
if sqlDB, _ := dropDB.DB(); sqlDB != nil {
|
||||
sqlDB.Close()
|
||||
}
|
||||
}()
|
||||
|
||||
if sqlDB, _ := dropDB.DB(); sqlDB != nil {
|
||||
sqlDB.SetMaxOpenConns(1)
|
||||
sqlDB.SetMaxIdleConns(0)
|
||||
sqlDB.SetConnMaxLifetime(time.Second)
|
||||
}
|
||||
|
||||
err = dropDB.Exec(fmt.Sprintf("DROP DATABASE IF EXISTS %s WITH (FORCE)", dbName)).Error
|
||||
|
||||
if err != nil {
|
||||
log.Errorf("failed to drop database %s: %v", dbName, err)
|
||||
}
|
||||
}
|
||||
|
||||
return replaceDBName(dsn, dbName), cleanup, nil
|
||||
}
|
||||
|
||||
func replaceDBName(dsn, newDBName string) string {
|
||||
re := regexp.MustCompile(`(?P<pre>[:/@])(?P<dbname>[^/?]+)(?P<post>\?|$)`)
|
||||
return re.ReplaceAllString(dsn, `${pre}`+newDBName+`${post}`)
|
||||
}
|
||||
@@ -18,6 +18,7 @@ import (
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map/controller/cache"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/peers/ephemeral"
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
"github.com/netbirdio/netbird/management/internals/server/config"
|
||||
"github.com/netbirdio/netbird/management/internals/shared/grpc"
|
||||
"github.com/netbirdio/netbird/management/server/account"
|
||||
@@ -30,6 +31,8 @@ import (
|
||||
"github.com/netbirdio/netbird/management/server/telemetry"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
sharedgrpc "github.com/netbirdio/netbird/shared/management/grpc"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
"github.com/netbirdio/netbird/util"
|
||||
@@ -61,6 +64,8 @@ type Controller struct {
|
||||
serverSupportedSyncMessageVersion sharedgrpc.SyncMessageVersion
|
||||
|
||||
perAccountServerSupportedSyncMessageVersions map[string]sharedgrpc.SyncMessageVersion
|
||||
|
||||
nmdataStore networkmapdb.NetworkMapDBStore
|
||||
}
|
||||
|
||||
type bufferUpdate struct {
|
||||
@@ -78,7 +83,7 @@ type bufferAffectedUpdate struct {
|
||||
|
||||
var _ network_map.Controller = (*Controller)(nil)
|
||||
|
||||
func NewController(ctx context.Context, store store.Store, metrics telemetry.AppMetrics, peersUpdateManager network_map.PeersUpdateManager, requestBuffer account.RequestBuffer, integratedPeerValidator integrated_validator.IntegratedValidator, settingsManager settings.Manager, dnsDomain string, proxyController port_forwarding.Controller, ephemeralPeersManager ephemeral.Manager, config *config.Config) *Controller {
|
||||
func NewController(ctx context.Context, store store.Store, metrics telemetry.AppMetrics, peersUpdateManager network_map.PeersUpdateManager, requestBuffer account.RequestBuffer, integratedPeerValidator integrated_validator.IntegratedValidator, settingsManager settings.Manager, dnsDomain string, proxyController port_forwarding.Controller, ephemeralPeersManager ephemeral.Manager, config *config.Config, nmdataStore networkmapdb.NetworkMapDBStore) *Controller {
|
||||
nMetrics, err := newMetrics(metrics.UpdateChannelMetrics())
|
||||
if err != nil {
|
||||
log.Fatal(fmt.Errorf("error creating metrics: %w", err))
|
||||
@@ -99,6 +104,7 @@ func NewController(ctx context.Context, store store.Store, metrics telemetry.App
|
||||
EphemeralPeersManager: ephemeralPeersManager,
|
||||
serverSupportedSyncMessageVersion: sharedgrpc.SyncMessageVersionFromConfig(config.HighestSupportedSyncMessageVersion),
|
||||
perAccountServerSupportedSyncMessageVersions: sharedgrpc.SyncMessageVersionsFromMap(config.PerAccountHighestSupportedSyncMessageVersion),
|
||||
nmdataStore: nmdataStore,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -147,6 +153,11 @@ func (c *Controller) CountStreams() int {
|
||||
|
||||
func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID string, reason types.UpdateReason) error {
|
||||
log.WithContext(ctx).Tracef("updating peers for account %s from %s", accountID, util.GetCallerName())
|
||||
|
||||
if nmData := c.getNetworkMapData(ctx, accountID); nmData != nil {
|
||||
return c.sendUpdateAccountPeersFromData(ctx, accountID, reason, nmData)
|
||||
}
|
||||
|
||||
account, err := c.requestBuffer.GetAccountWithBackpressure(ctx, accountID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get account: %v", err)
|
||||
@@ -167,7 +178,7 @@ func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID strin
|
||||
return nil
|
||||
}
|
||||
|
||||
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
|
||||
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get validate peers: %v", err)
|
||||
}
|
||||
@@ -254,7 +265,7 @@ func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID strin
|
||||
// proxyNetworkMap rides the envelope as a ProxyPatch sidecar;
|
||||
// the client merges it into Calculate()'s output the same
|
||||
// way the legacy server did via NetworkMap.Merge.
|
||||
update = grpc.ToComponentSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, p, nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, account.Settings, extraSetting, maps.Keys(peerGroups), dnsFwdPort)
|
||||
update = grpc.ToComponentSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, types.TwinPeer(p), nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, types.TwinAccountSettings(account.Settings), extraSetting, maps.Keys(peerGroups), dnsFwdPort)
|
||||
c.metrics.CountToComponentSyncResponseDuration(time.Since(start))
|
||||
|
||||
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
|
||||
@@ -275,7 +286,7 @@ func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID strin
|
||||
}
|
||||
|
||||
start = time.Now()
|
||||
update = grpc.ToSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, p, nil, nil, nmap, dnsDomain, postureChecks, dnsCache, account.Settings, extraSetting, maps.Keys(peerGroups), dnsFwdPort)
|
||||
update = grpc.ToSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, types.TwinPeer(p), nil, nil, nmap, dnsDomain, postureChecks, dnsCache, types.TwinAccountSettings(account.Settings), extraSetting, maps.Keys(peerGroups), dnsFwdPort)
|
||||
c.metrics.CountToSyncResponseDuration(time.Since(start))
|
||||
|
||||
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
|
||||
@@ -293,6 +304,247 @@ func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID strin
|
||||
return nil
|
||||
}
|
||||
|
||||
// sendUpdateAccountPeersFromData is the account-free variant of
|
||||
// sendUpdateAccountPeers: everything is computed from the network-map DB
|
||||
// store's twin data; only extra settings and validated peers are resolved at
|
||||
// runtime. Proxy network maps and policy injection, private-service zones,
|
||||
// group-to-user SSH mappings and forced routing-peer DNS resolution have no
|
||||
// DB-backed source yet and are omitted.
|
||||
func (c *Controller) sendUpdateAccountPeersFromData(ctx context.Context, accountID string, reason types.UpdateReason, nmData *networkmap.NetworkMapData) error {
|
||||
peersToUpdate := c.connectedPeersFromData(nmData, nil)
|
||||
if len(peersToUpdate) == 0 {
|
||||
return nil
|
||||
}
|
||||
return c.sendUpdatesFromData(ctx, accountID, nmData, peersToUpdate, &reason)
|
||||
}
|
||||
|
||||
// sendUpdateForAffectedPeersFromData is the account-free variant of
|
||||
// sendUpdateForAffectedPeers.
|
||||
func (c *Controller) sendUpdateForAffectedPeersFromData(ctx context.Context, accountID string, peerIDs []string, nmData *networkmap.NetworkMapData) error {
|
||||
affected := make(map[string]struct{}, len(peerIDs))
|
||||
for _, id := range peerIDs {
|
||||
affected[id] = struct{}{}
|
||||
}
|
||||
|
||||
peersToUpdate := c.connectedPeersFromData(nmData, affected)
|
||||
if len(peersToUpdate) == 0 {
|
||||
log.WithContext(ctx).Tracef("sendUpdateForAffectedPeersFromData: no peers to update (affected peers not found in data or no channels)")
|
||||
return nil
|
||||
}
|
||||
|
||||
log.WithContext(ctx).Tracef("sendUpdateForAffectedPeersFromData: sending network map to %d connected peers", len(peersToUpdate))
|
||||
|
||||
return c.sendUpdatesFromData(ctx, accountID, nmData, peersToUpdate, nil)
|
||||
}
|
||||
|
||||
func (c *Controller) connectedPeersFromData(nmData *networkmap.NetworkMapData, affected map[string]struct{}) []*nmdata.Peer {
|
||||
var result []*nmdata.Peer
|
||||
for _, peer := range nmData.Peers {
|
||||
if affected != nil {
|
||||
if _, ok := affected[peer.ID]; !ok {
|
||||
continue
|
||||
}
|
||||
}
|
||||
if c.peersUpdateManager.HasChannel(peer.ID) {
|
||||
result = append(result, peer)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func (c *Controller) sendUpdatesFromData(ctx context.Context, accountID string, nmData *networkmap.NetworkMapData, peersToUpdate []*nmdata.Peer, reason *types.UpdateReason) error {
|
||||
globalStart := time.Now()
|
||||
|
||||
extraSettings, err := c.settingsManager.GetExtraSettings(ctx, accountID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get flow enabled status: %v", err)
|
||||
}
|
||||
|
||||
dnsCache := &cache.DNSConfigCache{}
|
||||
dnsDomain := c.getDNSDomainFromData(nmData.AccountSettings)
|
||||
peersCustomZone := networkmap.PeersCustomZone(ctx, accountID, dnsDomain, nmData.Peers, ipv6AllowedPeersFromData(nmData))
|
||||
|
||||
dnsFwdPort := computeForwarderPortFromData(nmData.Peers, network_map.DnsForwarderPortMinVersion)
|
||||
|
||||
var wg sync.WaitGroup
|
||||
semaphore := make(chan struct{}, 10)
|
||||
|
||||
for _, peer := range peersToUpdate {
|
||||
if reason != nil && c.accountManagerMetrics != nil {
|
||||
c.accountManagerMetrics.CountNmapTriggered(string(reason.Resource), string(reason.Operation))
|
||||
}
|
||||
|
||||
wg.Add(1)
|
||||
semaphore <- struct{}{}
|
||||
go func(p *nmdata.Peer) {
|
||||
defer wg.Done()
|
||||
defer func() { <-semaphore }()
|
||||
|
||||
start := time.Now()
|
||||
|
||||
postureChecks := peerPostureChecksFromData(nmData, p.ID)
|
||||
|
||||
c.metrics.CountCalcPostureChecksDuration(time.Since(start))
|
||||
start = time.Now()
|
||||
|
||||
peerGroups := maps.Keys(nmData.GetPeerGroups(p.ID))
|
||||
var update *proto.SyncResponse
|
||||
|
||||
commonSyncMessageVersion := sharedgrpc.HighestCommonSyncMessageVersion(
|
||||
c.perAccountOrGlobalSupportedSyncMessageVersions(accountID),
|
||||
sharedgrpc.SyncMessageVersionFromConfig(&p.Meta.SyncMessageVersion))
|
||||
|
||||
log.WithContext(ctx).
|
||||
WithFields(log.Fields{
|
||||
"sync_message_version": commonSyncMessageVersion,
|
||||
"server_sync_message_version": c.perAccountOrGlobalSupportedSyncMessageVersions(accountID),
|
||||
"peer_sync_message_version": sharedgrpc.SyncMessageVersionFromConfig(&p.Meta.SyncMessageVersion),
|
||||
}).Debug("common highest sync message version")
|
||||
|
||||
if commonSyncMessageVersion == sharedgrpc.ComponentNetworkMap {
|
||||
components := nmData.GetPeerNetworkMapComponents(p.ID, peersCustomZone)
|
||||
|
||||
c.metrics.CountCalcPeerNetworkMapDuration(time.Since(start))
|
||||
|
||||
start = time.Now()
|
||||
update = grpc.ToComponentSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, p, nil, nil, components, nil, dnsDomain, postureChecks, nmData.AccountSettings, extraSettings, peerGroups, dnsFwdPort)
|
||||
c.metrics.CountToComponentSyncResponseDuration(time.Since(start))
|
||||
|
||||
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
|
||||
Update: update,
|
||||
MessageType: network_map.MessageTypeNetworkMap,
|
||||
})
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
nmap := networkMapFromData(ctx, nmData, p.ID, peersCustomZone)
|
||||
|
||||
c.metrics.CountCalcPeerNetworkMapDuration(time.Since(start))
|
||||
|
||||
start = time.Now()
|
||||
update = grpc.ToSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, p, nil, nil, nmap, dnsDomain, postureChecks, dnsCache, nmData.AccountSettings, extraSettings, peerGroups, dnsFwdPort)
|
||||
c.metrics.CountToSyncResponseDuration(time.Since(start))
|
||||
|
||||
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
|
||||
Update: update,
|
||||
MessageType: network_map.MessageTypeNetworkMap,
|
||||
})
|
||||
}(peer)
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
if c.accountManagerMetrics != nil {
|
||||
c.accountManagerMetrics.CountUpdateAccountPeersDuration(time.Since(globalStart))
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Controller) getNetworkMapData(ctx context.Context, accountID string) *networkmap.NetworkMapData {
|
||||
if c.nmdataStore == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
nmData, err := c.nmdataStore.GetNetworkMapData(ctx, accountID)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("failed to get network map data for account %s, falling back to account-based computation: %v", accountID, err)
|
||||
return nil
|
||||
}
|
||||
|
||||
return nmData
|
||||
}
|
||||
|
||||
func (c *Controller) getDNSDomainFromData(settings *nmdata.AccountSettingsInfo) string {
|
||||
if settings == nil || settings.DNSDomain == "" {
|
||||
return c.dnsDomain
|
||||
}
|
||||
return settings.DNSDomain
|
||||
}
|
||||
|
||||
func ipv6AllowedPeersFromData(nmData *networkmap.NetworkMapData) map[string]struct{} {
|
||||
result := make(map[string]struct{})
|
||||
if nmData.AccountSettings != nil {
|
||||
for _, groupID := range nmData.AccountSettings.IPv6EnabledGroups {
|
||||
group := nmData.Groups[groupID]
|
||||
if group == nil {
|
||||
continue
|
||||
}
|
||||
for _, peerID := range group.Peers {
|
||||
result[peerID] = struct{}{}
|
||||
}
|
||||
}
|
||||
}
|
||||
for id, p := range nmData.Peers {
|
||||
if p != nil && p.ProxyMeta.Embedded {
|
||||
result[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func networkMapFromData(ctx context.Context, nmData *networkmap.NetworkMapData, peerID string, peersCustomZone nmdata.CustomZone) *types.NetworkMap {
|
||||
components := nmData.GetPeerNetworkMapComponents(peerID, peersCustomZone)
|
||||
if components.IsEmpty() {
|
||||
return &types.NetworkMap{Network: components.Network}
|
||||
}
|
||||
return types.CalculateNetworkMapFromComponents(ctx, components)
|
||||
}
|
||||
|
||||
// peerPostureChecksFromData mirrors getPeerPostureChecks on the twin store. The
|
||||
// sync response only encodes process-check file paths, so only ProcessCheck is
|
||||
// converted back to the server posture type.
|
||||
func peerPostureChecksFromData(nmData *networkmap.NetworkMapData, peerID string) []*posture.Checks {
|
||||
if len(nmData.PostureChecks) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
peerPostureChecks := make(map[string]*posture.Checks)
|
||||
for _, policy := range nmData.Policies {
|
||||
if policy == nil || !policy.Enabled || len(policy.SourcePostureChecks) == 0 {
|
||||
continue
|
||||
}
|
||||
if !isPeerInPolicySourceGroupsFromData(nmData, peerID, policy) {
|
||||
continue
|
||||
}
|
||||
for _, checkID := range policy.SourcePostureChecks {
|
||||
twin := nmData.PostureChecks[checkID]
|
||||
if twin == nil {
|
||||
continue
|
||||
}
|
||||
peerPostureChecks[checkID] = postureChecksFromTwin(twin)
|
||||
}
|
||||
}
|
||||
|
||||
return maps.Values(peerPostureChecks)
|
||||
}
|
||||
|
||||
func isPeerInPolicySourceGroupsFromData(nmData *networkmap.NetworkMapData, peerID string, policy *nmdata.Policy) bool {
|
||||
for _, rule := range policy.Rules {
|
||||
if rule == nil || !rule.Enabled {
|
||||
continue
|
||||
}
|
||||
for _, groupID := range rule.Sources {
|
||||
if group := nmData.Groups[groupID]; group != nil && slices.Contains(group.Peers, peerID) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func postureChecksFromTwin(twin *nmdata.PostureChecks) *posture.Checks {
|
||||
checks := &posture.Checks{ID: twin.ID}
|
||||
if twin.Checks.ProcessCheck != nil {
|
||||
processes := make([]posture.Process, 0, len(twin.Checks.ProcessCheck.Processes))
|
||||
for _, p := range twin.Checks.ProcessCheck.Processes {
|
||||
processes = append(processes, posture.Process{LinuxPath: p.LinuxPath, MacPath: p.MacPath, WindowsPath: p.WindowsPath})
|
||||
}
|
||||
checks.Checks.ProcessCheck = &posture.ProcessCheck{Processes: processes}
|
||||
}
|
||||
return checks
|
||||
}
|
||||
|
||||
func (c *Controller) perAccountOrGlobalSupportedSyncMessageVersions(accountId string) sharedgrpc.SyncMessageVersion {
|
||||
if perAccount, ok := c.perAccountServerSupportedSyncMessageVersions[accountId]; ok {
|
||||
return perAccount
|
||||
@@ -325,6 +577,10 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s
|
||||
return nil
|
||||
}
|
||||
|
||||
if nmData := c.getNetworkMapData(ctx, accountID); nmData != nil {
|
||||
return c.sendUpdateForAffectedPeersFromData(ctx, accountID, peerIDs, nmData)
|
||||
}
|
||||
|
||||
account, err := c.requestBuffer.GetAccountWithBackpressure(ctx, accountID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get account: %v", err)
|
||||
@@ -340,7 +596,7 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s
|
||||
|
||||
log.WithContext(ctx).Tracef("sendUpdateForAffectedPeers: sending network map to %d connected peers", len(peersToUpdate))
|
||||
|
||||
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
|
||||
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get validate peers: %v", err)
|
||||
}
|
||||
@@ -426,7 +682,7 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s
|
||||
// proxyNetworkMap rides the envelope as a ProxyPatch sidecar;
|
||||
// the client merges it into Calculate()'s output the same
|
||||
// way the legacy server did via NetworkMap.Merge.
|
||||
update = grpc.ToComponentSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, p, nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, account.Settings, extraSetting, maps.Keys(peerGroups), dnsFwdPort)
|
||||
update = grpc.ToComponentSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, types.TwinPeer(p), nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, types.TwinAccountSettings(account.Settings), extraSetting, maps.Keys(peerGroups), dnsFwdPort)
|
||||
c.metrics.CountToComponentSyncResponseDuration(time.Since(start))
|
||||
|
||||
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
|
||||
@@ -447,7 +703,7 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s
|
||||
}
|
||||
|
||||
start = time.Now()
|
||||
update = grpc.ToSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, p, nil, nil, nmap, dnsDomain, postureChecks, dnsCache, account.Settings, extraSetting, maps.Keys(peerGroups), dnsFwdPort)
|
||||
update = grpc.ToSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, types.TwinPeer(p), nil, nil, nmap, dnsDomain, postureChecks, dnsCache, types.TwinAccountSettings(account.Settings), extraSetting, maps.Keys(peerGroups), dnsFwdPort)
|
||||
c.metrics.CountToSyncResponseDuration(time.Since(start))
|
||||
|
||||
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
|
||||
@@ -504,7 +760,7 @@ func (c *Controller) UpdateAccountPeer(ctx context.Context, accountId string, pe
|
||||
return fmt.Errorf("peer %s doesn't exists in account %s", peerId, accountId)
|
||||
}
|
||||
|
||||
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
|
||||
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get validated peers: %v", err)
|
||||
}
|
||||
@@ -564,7 +820,7 @@ func (c *Controller) UpdateAccountPeer(ctx context.Context, accountId string, pe
|
||||
// proxyNetworkMap rides the envelope as a ProxyPatch sidecar;
|
||||
// the client merges it into Calculate()'s output the same
|
||||
// way the legacy server did via NetworkMap.Merge.
|
||||
update = grpc.ToComponentSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, peer, nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, account.Settings, extraSettings, maps.Keys(peerGroups), dnsFwdPort)
|
||||
update = grpc.ToComponentSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, types.TwinPeer(peer), nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, types.TwinAccountSettings(account.Settings), extraSettings, maps.Keys(peerGroups), dnsFwdPort)
|
||||
|
||||
c.peersUpdateManager.SendUpdate(ctx, peer.ID, &network_map.UpdateMessage{
|
||||
Update: update,
|
||||
@@ -581,7 +837,7 @@ func (c *Controller) UpdateAccountPeer(ctx context.Context, accountId string, pe
|
||||
nmap.Merge(proxyNetworkMap)
|
||||
}
|
||||
|
||||
update = grpc.ToSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, peer, nil, nil, nmap, dnsDomain, postureChecks, dnsCache, account.Settings, extraSettings, maps.Keys(peerGroups), dnsFwdPort)
|
||||
update = grpc.ToSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, types.TwinPeer(peer), nil, nil, nmap, dnsDomain, postureChecks, dnsCache, types.TwinAccountSettings(account.Settings), extraSettings, maps.Keys(peerGroups), dnsFwdPort)
|
||||
|
||||
c.peersUpdateManager.SendUpdate(ctx, peer.ID, &network_map.UpdateMessage{
|
||||
Update: update,
|
||||
@@ -641,7 +897,7 @@ func (c *Controller) GetValidatedPeerWithComponents(ctx context.Context, isRequi
|
||||
if err != nil {
|
||||
return nil, nil, nil, nil, 0, err
|
||||
}
|
||||
return peer, &types.NetworkMapComponents{Network: network.Copy()}, nil, nil, 0, nil
|
||||
return peer, &types.NetworkMapComponents{Network: types.TwinNetwork(network)}, nil, nil, 0, nil
|
||||
}
|
||||
|
||||
account, err := c.requestBuffer.GetAccountWithBackpressure(ctx, accountID)
|
||||
@@ -651,7 +907,7 @@ func (c *Controller) GetValidatedPeerWithComponents(ctx context.Context, isRequi
|
||||
|
||||
c.injectAllProxyPolicies(ctx, account)
|
||||
|
||||
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
|
||||
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
|
||||
if err != nil {
|
||||
return nil, nil, nil, nil, 0, err
|
||||
}
|
||||
@@ -794,7 +1050,7 @@ func (c *Controller) GetValidatedPeerWithMap(ctx context.Context, isRequiresAppr
|
||||
}
|
||||
|
||||
emptyMap := &types.NetworkMap{
|
||||
Network: network.Copy(),
|
||||
Network: types.TwinNetwork(network),
|
||||
}
|
||||
return emptyMap, nil, 0, nil
|
||||
}
|
||||
@@ -806,7 +1062,7 @@ func (c *Controller) GetValidatedPeerWithMap(ctx context.Context, isRequiresAppr
|
||||
|
||||
c.injectAllProxyPolicies(ctx, account)
|
||||
|
||||
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
|
||||
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
|
||||
if err != nil {
|
||||
return nil, nil, 0, err
|
||||
}
|
||||
@@ -908,20 +1164,36 @@ func (c *Controller) StartWarmup(ctx context.Context) {
|
||||
// computeForwarderPort checks if all peers in the account have updated to a specific version or newer.
|
||||
// If all peers have the required version, it returns the new well-known port (22054), otherwise returns 0.
|
||||
func computeForwarderPort(peers []*nbpeer.Peer, requiredVersion string) int64 {
|
||||
if len(peers) == 0 {
|
||||
versions := make([]string, 0, len(peers))
|
||||
for _, peer := range peers {
|
||||
versions = append(versions, peer.Meta.WtVersion)
|
||||
}
|
||||
return computeForwarderPortFromVersions(versions, requiredVersion)
|
||||
}
|
||||
|
||||
func computeForwarderPortFromData(peers map[string]*nmdata.Peer, requiredVersion string) int64 {
|
||||
versions := make([]string, 0, len(peers))
|
||||
for _, peer := range peers {
|
||||
versions = append(versions, peer.Meta.WtVersion)
|
||||
}
|
||||
return computeForwarderPortFromVersions(versions, requiredVersion)
|
||||
}
|
||||
|
||||
func computeForwarderPortFromVersions(wtVersions []string, requiredVersion string) int64 {
|
||||
if len(wtVersions) == 0 {
|
||||
return int64(network_map.OldForwarderPort)
|
||||
}
|
||||
|
||||
reqVer := semver.Canonical(requiredVersion)
|
||||
|
||||
// Check if all peers have the required version or newer
|
||||
for _, peer := range peers {
|
||||
for _, wtVersion := range wtVersions {
|
||||
|
||||
// Development version is always supported
|
||||
if version.IsDevelopmentVersion(peer.Meta.WtVersion) {
|
||||
if version.IsDevelopmentVersion(wtVersion) {
|
||||
continue
|
||||
}
|
||||
peerVersion := semver.Canonical("v" + peer.Meta.WtVersion)
|
||||
peerVersion := semver.Canonical("v" + wtVersion)
|
||||
if peerVersion == "" {
|
||||
// If any peer doesn't have version info, return 0
|
||||
return int64(network_map.OldForwarderPort)
|
||||
@@ -1055,7 +1327,7 @@ func (c *Controller) GetNetworkMap(ctx context.Context, peerID string) (*types.N
|
||||
groups[groupID] = group.Peers
|
||||
}
|
||||
|
||||
validatedPeers, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
|
||||
validatedPeers, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
135
management/internals/network_map_db/db_store.go
Normal file
135
management/internals/network_map_db/db_store.go
Normal file
@@ -0,0 +1,135 @@
|
||||
package networkmapdb
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"reflect"
|
||||
"strings"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/integrations/integrated_validator"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/rs/xid"
|
||||
)
|
||||
|
||||
const (
|
||||
NMAP_STRUCT_TAG = "nmap"
|
||||
NMAP_SKIP = "skip"
|
||||
NMAP_MAP_TO = "map_to"
|
||||
)
|
||||
|
||||
type NetworkMapDBStore interface {
|
||||
GetGroups(ctx context.Context, accountId string) ([]nmdata.Group, map[string]map[string]any, error)
|
||||
GetPeers(ctx context.Context, accountId string) ([]nmdata.Peer, map[string]*nmdata.Peer, error)
|
||||
GetPolicies(ctx context.Context, accountId string) ([]nmdata.Policy, map[string]map[string]any, map[string]map[string]any, error)
|
||||
GetRoutes(ctx context.Context, accountId string) ([]nmdata.Route, error)
|
||||
GetNameServerGroups(ctx context.Context, accountId string) ([]nmdata.NameServerGroup, error)
|
||||
GetNetworkResources(ctx context.Context, accountId string) ([]nmdata.NetworkResource, error)
|
||||
GetNetworkRouters(ctx context.Context, accountId string) (map[string]map[string]*nmdata.NetworkRouter, error)
|
||||
GetNetwork(ctx context.Context, accountId string) (nmdata.Network, error)
|
||||
GetAppliedZoneCandidates(ctx context.Context, accountId string) ([]networkmap.AppliedZoneCandidate, error)
|
||||
GetAccountSettings(ctx context.Context, accountId string) (nmdata.AccountSettingsInfo, error)
|
||||
GetPostureChecks(ctx context.Context, accountId string) ([]nmdata.PostureChecks, error)
|
||||
GetNetworkMapData(ctx context.Context, accountId string) (*networkmap.NetworkMapData, error)
|
||||
GetAllowedUsers(ctx context.Context, accountId string) (map[string]struct{}, map[string][]string, error)
|
||||
GetDnsSettings(ctx context.Context, accountId string) (nmdata.DNSSettings, error)
|
||||
}
|
||||
|
||||
type NetworkMapDBStoreImpl struct {
|
||||
store NetworkMapDBStore
|
||||
integratedPeerValidator integrated_validator.IntegratedValidator
|
||||
}
|
||||
|
||||
func FromSqlTypesToSharedTypes(src reflect.Value, dst reflect.Value) error {
|
||||
typ := src.Elem().Type()
|
||||
|
||||
for i := 0; i < typ.NumField(); i++ {
|
||||
f := typ.Field(i)
|
||||
|
||||
fieldTags := make(map[string]string)
|
||||
if v := f.Tag.Get(NMAP_STRUCT_TAG); v != "" {
|
||||
for _, t := range strings.Split(v, ",") {
|
||||
kv := tagFromString(t)
|
||||
fieldTags[kv.Key] = kv.Value
|
||||
}
|
||||
}
|
||||
if _, ok := fieldTags[NMAP_SKIP]; ok {
|
||||
continue
|
||||
}
|
||||
if f.PkgPath != "" { // skip unexported fields
|
||||
continue
|
||||
}
|
||||
dstFieldName := f.Name
|
||||
if override, ok := fieldTags[NMAP_MAP_TO]; ok {
|
||||
dstFieldName = override
|
||||
}
|
||||
|
||||
dstField := dst.Elem().FieldByName(dstFieldName)
|
||||
if !dstField.IsValid() {
|
||||
return errors.New("unsupported type in destination field: " + dstFieldName)
|
||||
}
|
||||
|
||||
srcField := src.Elem().Field(i)
|
||||
srcFieldType := srcField.Type().String()
|
||||
switch srcFieldType {
|
||||
case "string":
|
||||
s := srcField.Interface().(string)
|
||||
dstField.SetString(s)
|
||||
case "sql.NullString":
|
||||
s := srcField.Interface().(sql.NullString)
|
||||
if s.Valid {
|
||||
dstField.SetString(s.String)
|
||||
}
|
||||
if (dstFieldName == "PublicId" || dstFieldName == "PublicID") && s.String == "" {
|
||||
dstField.SetString(xid.New().String()) // TODO (dmitri) this needs to be removed to support delta updates
|
||||
}
|
||||
case "sql.NullTime":
|
||||
s := srcField.Interface().(sql.NullTime)
|
||||
if s.Valid {
|
||||
if dstField.Kind() == reflect.Ptr {
|
||||
t := reflect.ValueOf(&s.Time).Elem()
|
||||
dstField.Set(t.Addr())
|
||||
} else {
|
||||
dstField.Set(reflect.ValueOf(s.Time))
|
||||
}
|
||||
}
|
||||
case "sql.NullBool":
|
||||
s := srcField.Interface().(sql.NullBool)
|
||||
if s.Valid {
|
||||
dstField.SetBool(s.Bool)
|
||||
}
|
||||
case "sql.NullInt64":
|
||||
s := srcField.Interface().(sql.NullInt64)
|
||||
if s.Valid {
|
||||
dstField.SetInt(s.Int64)
|
||||
}
|
||||
case "json.RawMessage":
|
||||
s := srcField.Interface().(json.RawMessage)
|
||||
json.Unmarshal(s, dstField.Addr().Interface())
|
||||
case "[]string":
|
||||
if srcField.IsNil() {
|
||||
return nil
|
||||
}
|
||||
dstv := reflect.MakeSlice(dstField.Type(), srcField.Len(), srcField.Cap())
|
||||
reflect.Copy(dstv, srcField)
|
||||
dstField.Set(dstv)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
type fieldTag struct {
|
||||
Key string
|
||||
Value string
|
||||
}
|
||||
|
||||
func tagFromString(t string) fieldTag {
|
||||
kv := strings.Split(t, ":")
|
||||
if len(kv) == 1 {
|
||||
return fieldTag{Key: strings.TrimSpace(kv[0])}
|
||||
}
|
||||
return fieldTag{Key: strings.TrimSpace(kv[0]), Value: strings.TrimSpace(kv[1])}
|
||||
}
|
||||
@@ -0,0 +1,84 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
const (
|
||||
GetAccountSettingsQuery = `
|
||||
select settings_peer_login_expiration_enabled as peer_login_expiration_enabled,
|
||||
settings_peer_login_expiration as peer_login_expiration,
|
||||
settings_peer_inactivity_expiration_enabled as peer_inactivity_expiration_enabled,
|
||||
settings_peer_inactivity_expiration as peer_inactivity_expiration,
|
||||
settings_dns_domain as dns_domain,
|
||||
settings_ipv6_enabled_groups as ipv6_enabled_groups,
|
||||
settings_routing_peer_dns_resolution_enabled as routing_peer_dns_resolution_enabled,
|
||||
settings_lazy_connection_enabled as lazy_connection_enabled,
|
||||
settings_auto_update_version as auto_update_version,
|
||||
settings_auto_update_always as auto_update_always,
|
||||
settings_metrics_push_enabled as metrics_push_enabled
|
||||
from accounts
|
||||
where id=$1
|
||||
`
|
||||
)
|
||||
|
||||
func (pg *PgStore) GetAccountSettings(ctx context.Context, accountId string) (nmdata.AccountSettingsInfo, error) {
|
||||
c, err := pg.Pool.Acquire(ctx)
|
||||
if err != nil {
|
||||
return nmdata.AccountSettingsInfo{}, err
|
||||
}
|
||||
return GetAccountSettingsViaPgxConnection(ctx, c.Conn(), accountId)
|
||||
|
||||
}
|
||||
|
||||
func GetAccountSettingsViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (nmdata.AccountSettingsInfo, error) {
|
||||
rows, err := con.Query(ctx, GetAccountSettingsQuery, accountId)
|
||||
if err != nil {
|
||||
return nmdata.AccountSettingsInfo{}, err
|
||||
}
|
||||
|
||||
settings, err := pgx.CollectOneRow(rows, pgx.RowToStructByName[account])
|
||||
if err != nil {
|
||||
return nmdata.AccountSettingsInfo{}, err
|
||||
}
|
||||
|
||||
settingsInfo := nmdata.AccountSettingsInfo{
|
||||
PeerLoginExpirationEnabled: settings.PeerLoginExpirationEnabled.Bool,
|
||||
PeerLoginExpiration: time.Duration(settings.PeerLoginExpiration.Int64),
|
||||
PeerInactivityExpirationEnabled: settings.PeerInactivityExpirationEnabled.Bool,
|
||||
PeerInactivityExpiration: time.Duration(settings.PeerInactivityExpiration.Int64),
|
||||
DNSDomain: settings.DNSDomain.String,
|
||||
RoutingPeerDNSResolutionEnabled: settings.RoutingPeerDNSResolutionEnabled.Bool,
|
||||
LazyConnectionEnabled: settings.LazyConnectionEnabled.Bool,
|
||||
AutoUpdateVersion: settings.AutoUpdateVersion.String,
|
||||
AutoUpdateAlways: settings.AutoUpdateAlways.Bool,
|
||||
MetricsPushEnabled: settings.MetricsPushEnabled.Bool,
|
||||
}
|
||||
if settings.IPv6EnabledGroups != nil {
|
||||
if err := json.Unmarshal(settings.IPv6EnabledGroups, &settingsInfo.IPv6EnabledGroups); err != nil {
|
||||
return nmdata.AccountSettingsInfo{}, err
|
||||
}
|
||||
}
|
||||
|
||||
return settingsInfo, nil
|
||||
}
|
||||
|
||||
type account struct {
|
||||
PeerLoginExpirationEnabled sql.NullBool
|
||||
PeerLoginExpiration sql.NullInt64
|
||||
PeerInactivityExpirationEnabled sql.NullBool
|
||||
PeerInactivityExpiration sql.NullInt64
|
||||
DNSDomain sql.NullString
|
||||
IPv6EnabledGroups json.RawMessage
|
||||
RoutingPeerDNSResolutionEnabled sql.NullBool
|
||||
LazyConnectionEnabled sql.NullBool
|
||||
AutoUpdateVersion sql.NullString
|
||||
AutoUpdateAlways sql.NullBool
|
||||
MetricsPushEnabled sql.NullBool
|
||||
}
|
||||
125
management/internals/network_map_db/pgsql/dns.go
Normal file
125
management/internals/network_map_db/pgsql/dns.go
Normal file
@@ -0,0 +1,125 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"reflect"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/miekg/dns"
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
var DnsUnsupportedRecordTypeError = errors.New("unsupported record type")
|
||||
|
||||
const (
|
||||
GetAccountZonesQuery = `
|
||||
select zones.id as id, domain, enable_search_domain as search_domain_disabled, distribution_groups,
|
||||
r.name as record_name, r.type as record_type, 'IN' record_class, r.ttl as record_ttl, r.content as record_rdata
|
||||
from zones
|
||||
left join records as r on r.zone_id = zones.id
|
||||
where zones.account_id=$1
|
||||
`
|
||||
)
|
||||
|
||||
func (pg *PgStore) GetAppliedZoneCandidates(ctx context.Context, accountId string) ([]networkmap.AppliedZoneCandidate, error) {
|
||||
c, err := pg.Pool.Acquire(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetAppliedZoneCandidatesViaPgxConnection(ctx, c.Conn(), accountId)
|
||||
}
|
||||
|
||||
func GetAppliedZoneCandidatesViaPgxConnection(ctx context.Context, conn *pgx.Conn, accountId string) ([]networkmap.AppliedZoneCandidate, error) {
|
||||
rows, err := conn.Query(ctx, GetAccountZonesQuery, accountId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
zones, err := pgx.CollectRows(rows, pgx.RowToStructByName[zone])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
toret := make([]networkmap.AppliedZoneCandidate, 0, len(zones))
|
||||
currentZoneId := ""
|
||||
for _, z := range zones {
|
||||
zone := nmdata.CustomZone{}
|
||||
err := networkmapdb.FromSqlTypesToSharedTypes(
|
||||
reflect.ValueOf(&z), reflect.ValueOf(&zone))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var distributionGroups []string
|
||||
if err := json.Unmarshal(z.DistributionGroups, &distributionGroups); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
rtype, rdata, err := recordTypeAndRdata(z.RecordType.String, z.RecordRData.String)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
record := nmdata.SimpleRecord{
|
||||
Name: z.RecordName.String,
|
||||
Class: z.RecordClass.String,
|
||||
TTL: int(z.RecordTTL.Int64),
|
||||
RData: rdata,
|
||||
Type: rtype,
|
||||
}
|
||||
zone.Records = []nmdata.SimpleRecord{record}
|
||||
|
||||
if len(toret) == 0 {
|
||||
toret = append(toret, appliedZoneCandidateFromZone(zone, distributionGroups))
|
||||
currentZoneId = z.Id
|
||||
continue
|
||||
}
|
||||
|
||||
if z.Id == currentZoneId {
|
||||
lastZone := &toret[len(toret)-1]
|
||||
lastZone.Zone.Records = append(lastZone.Zone.Records, record)
|
||||
continue
|
||||
}
|
||||
|
||||
toret = append(toret, appliedZoneCandidateFromZone(zone, distributionGroups))
|
||||
currentZoneId = z.Id
|
||||
}
|
||||
return toret, nil
|
||||
}
|
||||
|
||||
type zone struct {
|
||||
Id string `nmap:"skip"`
|
||||
DistributionGroups json.RawMessage `nmap:"skip"`
|
||||
Domain sql.NullString
|
||||
SearchDomainDisabled sql.NullBool
|
||||
RecordName sql.NullString `nmap:"skip"`
|
||||
RecordType sql.NullString `nmap:"skip"`
|
||||
RecordClass sql.NullString `nmap:"skip"`
|
||||
RecordTTL sql.NullInt64 `nmap:"skip"`
|
||||
RecordRData sql.NullString `nmap:"skip"`
|
||||
}
|
||||
|
||||
func recordTypeAndRdata(t, rdata string) (int, string, error) {
|
||||
switch t {
|
||||
case "A":
|
||||
return int(dns.TypeA), rdata, nil
|
||||
case "AAAA":
|
||||
return int(dns.TypeAAAA), rdata, nil
|
||||
case "CNAME":
|
||||
return int(dns.TypeCNAME), dns.Fqdn(rdata), nil
|
||||
default:
|
||||
return 0, "", fmt.Errorf("record type: %s %w", t, DnsUnsupportedRecordTypeError)
|
||||
}
|
||||
}
|
||||
|
||||
func appliedZoneCandidateFromZone(z nmdata.CustomZone, distributionGroups []string) networkmap.AppliedZoneCandidate {
|
||||
return networkmap.AppliedZoneCandidate{
|
||||
DistributionGroups: distributionGroups,
|
||||
Zone: z,
|
||||
}
|
||||
}
|
||||
49
management/internals/network_map_db/pgsql/dns_settings.go
Normal file
49
management/internals/network_map_db/pgsql/dns_settings.go
Normal file
@@ -0,0 +1,49 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
const (
|
||||
GetDnsSettingsQuery = `
|
||||
select dns_settings_disabled_management_groups
|
||||
from accounts
|
||||
where id=$1
|
||||
`
|
||||
)
|
||||
|
||||
func (pg *PgStore) GetDnsSettings(ctx context.Context, accountId string) (nmdata.DNSSettings, error) {
|
||||
c, err := pg.Pool.Acquire(ctx)
|
||||
if err != nil {
|
||||
return nmdata.DNSSettings{}, err
|
||||
}
|
||||
return GetDnsSettingsViaPgxConnection(ctx, c.Conn(), accountId)
|
||||
}
|
||||
|
||||
func GetDnsSettingsViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (nmdata.DNSSettings, error) {
|
||||
rows, err := con.Query(ctx, GetDnsSettingsQuery, accountId)
|
||||
if err != nil {
|
||||
return nmdata.DNSSettings{}, err
|
||||
}
|
||||
|
||||
return pgx.CollectOneRow(rows, rowToDnsSettings)
|
||||
}
|
||||
|
||||
func rowToDnsSettings(row pgx.CollectableRow) (nmdata.DNSSettings, error) {
|
||||
var value nmdata.DNSSettings
|
||||
var settings json.RawMessage
|
||||
|
||||
if err := row.Scan(&settings); err != nil {
|
||||
return value, err
|
||||
}
|
||||
|
||||
if err := json.Unmarshal(settings, &value.DisabledManagementGroups); err != nil {
|
||||
return value, err
|
||||
}
|
||||
|
||||
return value, nil
|
||||
}
|
||||
38
management/internals/network_map_db/pgsql/dns_test.go
Normal file
38
management/internals/network_map_db/pgsql/dns_test.go
Normal file
@@ -0,0 +1,38 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestRecordTypeAndRdata(t *testing.T) {
|
||||
var tests = []struct {
|
||||
recordType string
|
||||
expectedRecordType int
|
||||
rdata string
|
||||
expectedRdata string
|
||||
expectedErr error
|
||||
}{
|
||||
{recordType: "A", expectedRecordType: 1, rdata: "test.com", expectedRdata: "test.com", expectedErr: nil},
|
||||
{recordType: "AAAA", expectedRecordType: 28, rdata: "test.com", expectedRdata: "test.com", expectedErr: nil},
|
||||
{recordType: "CNAME", expectedRecordType: 5, rdata: "test.com", expectedRdata: "test.com.", expectedErr: nil},
|
||||
{recordType: "CNAME", expectedRecordType: 5, rdata: "test.com.", expectedRdata: "test.com.", expectedErr: nil},
|
||||
{recordType: "TypeMX", expectedErr: DnsUnsupportedRecordTypeError},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.recordType, func(t *testing.T) {
|
||||
recordType, rdata, err := recordTypeAndRdata(tt.recordType, tt.rdata)
|
||||
|
||||
if tt.expectedErr != nil {
|
||||
assert.ErrorIs(t, err, DnsUnsupportedRecordTypeError)
|
||||
return
|
||||
}
|
||||
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, recordType, tt.expectedRecordType)
|
||||
assert.Equal(t, rdata, tt.expectedRdata)
|
||||
})
|
||||
}
|
||||
}
|
||||
72
management/internals/network_map_db/pgsql/group.go
Normal file
72
management/internals/network_map_db/pgsql/group.go
Normal file
@@ -0,0 +1,72 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"reflect"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
const (
|
||||
GetGroupsQuery = `
|
||||
select id, name, public_id, resources,
|
||||
(
|
||||
select array_agg(group_peers.peer_id)
|
||||
from group_peers
|
||||
where group_peers.group_id = groups.id
|
||||
) as peers
|
||||
from groups where account_id=$1
|
||||
`
|
||||
)
|
||||
|
||||
// we also return a resource-to-group index.
|
||||
// an alternative is to add json indexes, query this directly. Not sure how expensive
|
||||
// json indexes are. TODO (dmitri) verify and maybe change the implementation here.
|
||||
func (pg *PgStore) GetGroups(ctx context.Context, accountId string) ([]nmdata.Group, map[string]map[string]any, error) {
|
||||
c, err := pg.Pool.Acquire(ctx)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
return GetGroupsViaPgxConnection(ctx, c.Conn(), accountId)
|
||||
}
|
||||
|
||||
func GetGroupsViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.Group, map[string]map[string]any, error) {
|
||||
rows, err := con.Query(ctx, GetGroupsQuery, accountId)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
groups, err := pgx.CollectRows(rows, pgx.RowToStructByName[group])
|
||||
toret := make([]nmdata.Group, 0, len(groups))
|
||||
resourceToGroupIdx := make(map[string]map[string]any)
|
||||
|
||||
for _, g := range groups {
|
||||
dg := nmdata.Group{}
|
||||
err := networkmapdb.FromSqlTypesToSharedTypes(
|
||||
reflect.ValueOf(&g), reflect.ValueOf(&dg))
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
toret = append(toret, dg)
|
||||
for _, resource := range dg.Resources {
|
||||
if _, ok := resourceToGroupIdx[resource.ID]; !ok {
|
||||
resourceToGroupIdx[resource.ID] = make(map[string]any)
|
||||
}
|
||||
resourceToGroupIdx[resource.ID][g.ID] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
return toret, resourceToGroupIdx, err
|
||||
}
|
||||
|
||||
type group struct {
|
||||
ID string
|
||||
Name sql.NullString
|
||||
PublicID sql.NullString
|
||||
Resources json.RawMessage
|
||||
Peers []string
|
||||
}
|
||||
283
management/internals/network_map_db/pgsql/group_test.go
Normal file
283
management/internals/network_map_db/pgsql/group_test.go
Normal file
@@ -0,0 +1,283 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
_ "embed"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestGetGroups(t *testing.T) {
|
||||
ctx := context.TODO()
|
||||
|
||||
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
|
||||
assert.NoError(t, err)
|
||||
// err = loadSQL(ctx, s.pool, initDb)
|
||||
//assert.NoError(t, err)
|
||||
|
||||
_, err = s.Pool.Query(ctx, "insert into groups (id, account_id, name, resources, public_id) VALUES('test-group-id-1','ck7bnf2t2r9s739pkug0','test-group-1', '[{\"ID\":\"cui7q2jl0ubs73d8qpi0\",\"Type\":\"host\"}]','public-id-1')")
|
||||
assert.NoError(t, err)
|
||||
|
||||
groups, _, err := s.GetGroups(ctx, "ck7bnf2t2r9s739pkug0") //"ckd7ee2fic3c73dtendg")
|
||||
assert.NoError(t, err)
|
||||
assert.Contains(t,
|
||||
groups,
|
||||
nmdata.Group{Name: "test-group-1", PublicID: "public-id-1", Resources: []nmdata.Resource{{ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
)
|
||||
assert.Contains(t,
|
||||
groups,
|
||||
nmdata.Group{Name: "All", PublicID: "d9aejspvcsu517nkh4a0", Resources: []nmdata.Resource{{ID: "cui7olrl0ubs73d8qpe0", Type: "subnet"}, {ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
)
|
||||
}
|
||||
|
||||
func TestGetPeers(t *testing.T) {
|
||||
ctx := context.TODO()
|
||||
|
||||
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
|
||||
assert.NoError(t, err)
|
||||
// err = loadSQL(ctx, s.pool, initDb)
|
||||
//assert.NoError(t, err)
|
||||
|
||||
peers, clusterToPeerIdx, err := s.GetPeers(ctx, "d8pqjvbl0ubs73e8cjkg") //"ckd7ee2fic3c73dtendg")
|
||||
assert.NoError(t, err)
|
||||
assert.NotEmpty(t, clusterToPeerIdx)
|
||||
|
||||
fmt.Print(peers)
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "test-group-1", PublicID: "public-id-1", Resources: []nmdata.Resource{{ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "All", PublicID: "d9aejspvcsu517nkh4a0", Resources: []nmdata.Resource{{ID: "cui7olrl0ubs73d8qpe0", Type: "subnet"}, {ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
}
|
||||
|
||||
func TestGetPolicies(t *testing.T) {
|
||||
ctx := context.TODO()
|
||||
|
||||
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
|
||||
assert.NoError(t, err)
|
||||
// err = loadSQL(ctx, s.pool, initDb)
|
||||
//assert.NoError(t, err)
|
||||
|
||||
peers, idx1, idx2, err := s.GetPolicies(ctx, "ck7bnf2t2r9s739pkug0") //"ckd7ee2fic3c73dtendg")
|
||||
assert.NoError(t, err)
|
||||
assert.NotEmpty(t, idx1)
|
||||
assert.NotEmpty(t, idx2)
|
||||
|
||||
fmt.Print(peers)
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "test-group-1", PublicID: "public-id-1", Resources: []nmdata.Resource{{ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "All", PublicID: "d9aejspvcsu517nkh4a0", Resources: []nmdata.Resource{{ID: "cui7olrl0ubs73d8qpe0", Type: "subnet"}, {ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
}
|
||||
|
||||
func TestGetRoutes(t *testing.T) {
|
||||
ctx := context.TODO()
|
||||
|
||||
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
|
||||
assert.NoError(t, err)
|
||||
// err = loadSQL(ctx, s.pool, initDb)
|
||||
//assert.NoError(t, err)
|
||||
|
||||
peers, err := s.GetRoutes(ctx, "csg5iabl0ubs7398nf1g") //"ckd7ee2fic3c73dtendg")
|
||||
assert.NoError(t, err)
|
||||
|
||||
fmt.Print(peers)
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "test-group-1", PublicID: "public-id-1", Resources: []nmdata.Resource{{ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "All", PublicID: "d9aejspvcsu517nkh4a0", Resources: []nmdata.Resource{{ID: "cui7olrl0ubs73d8qpe0", Type: "subnet"}, {ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
}
|
||||
|
||||
func TestGetNSGroups(t *testing.T) {
|
||||
ctx := context.TODO()
|
||||
|
||||
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
|
||||
assert.NoError(t, err)
|
||||
// err = loadSQL(ctx, s.pool, initDb)
|
||||
//assert.NoError(t, err)
|
||||
|
||||
groups, err := s.GetNameServerGroups(ctx, "cl3h77qfic3c738mkja0") //"ckd7ee2fic3c73dtendg")
|
||||
assert.NoError(t, err)
|
||||
|
||||
fmt.Print(groups)
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "test-group-1", PublicID: "public-id-1", Resources: []nmdata.Resource{{ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "All", PublicID: "d9aejspvcsu517nkh4a0", Resources: []nmdata.Resource{{ID: "cui7olrl0ubs73d8qpe0", Type: "subnet"}, {ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
}
|
||||
|
||||
func TestGetNetworkResources(t *testing.T) {
|
||||
ctx := context.TODO()
|
||||
|
||||
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
|
||||
assert.NoError(t, err)
|
||||
// err = loadSQL(ctx, s.pool, initDb)
|
||||
//assert.NoError(t, err)
|
||||
|
||||
res, err := s.GetNetworkResources(ctx, "cag86v2t2r9s73d0416g") //"ckd7ee2fic3c73dtendg")
|
||||
assert.NoError(t, err)
|
||||
|
||||
fmt.Print(res)
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "test-group-1", PublicID: "public-id-1", Resources: []nmdata.Resource{{ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "All", PublicID: "d9aejspvcsu517nkh4a0", Resources: []nmdata.Resource{{ID: "cui7olrl0ubs73d8qpe0", Type: "subnet"}, {ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
}
|
||||
|
||||
func TestGetNetworkRouters(t *testing.T) {
|
||||
ctx := context.TODO()
|
||||
|
||||
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
|
||||
assert.NoError(t, err)
|
||||
// err = loadSQL(ctx, s.pool, initDb)
|
||||
//assert.NoError(t, err)
|
||||
|
||||
c, _ := s.Pool.Acquire(ctx)
|
||||
res, err := GetNetworkRoutersViaPgxConnection(ctx, c.Conn(), "d29f99jl0ubs73cm8ce0") //"ckd7ee2fic3c73dtendg")
|
||||
assert.NoError(t, err)
|
||||
|
||||
fmt.Print(res)
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "test-group-1", PublicID: "public-id-1", Resources: []nmdata.Resource{{ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "All", PublicID: "d9aejspvcsu517nkh4a0", Resources: []nmdata.Resource{{ID: "cui7olrl0ubs73d8qpe0", Type: "subnet"}, {ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
}
|
||||
|
||||
func TestGetNetwork(t *testing.T) {
|
||||
ctx := context.TODO()
|
||||
|
||||
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
|
||||
assert.NoError(t, err)
|
||||
// err = loadSQL(ctx, s.pool, initDb)
|
||||
//assert.NoError(t, err)
|
||||
|
||||
n, err := s.GetNetwork(ctx, "d29f99jl0ubs73cm8ce0") //"ckd7ee2fic3c73dtendg")
|
||||
assert.NoError(t, err)
|
||||
|
||||
fmt.Print(n)
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "test-group-1", PublicID: "public-id-1", Resources: []nmdata.Resource{{ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "All", PublicID: "d9aejspvcsu517nkh4a0", Resources: []nmdata.Resource{{ID: "cui7olrl0ubs73d8qpe0", Type: "subnet"}, {ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
}
|
||||
|
||||
func TestGetAccountZones(t *testing.T) {
|
||||
ctx := context.TODO()
|
||||
|
||||
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
|
||||
assert.NoError(t, err)
|
||||
// err = loadSQL(ctx, s.pool, initDb)
|
||||
//assert.NoError(t, err)
|
||||
|
||||
zones, err := s.GetAppliedZoneCandidates(ctx, "d4g66rjl0ubs73b2q3b0") //"ckd7ee2fic3c73dtendg")
|
||||
assert.NoError(t, err)
|
||||
|
||||
fmt.Print(zones)
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "test-group-1", PublicID: "public-id-1", Resources: []nmdata.Resource{{ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "All", PublicID: "d9aejspvcsu517nkh4a0", Resources: []nmdata.Resource{{ID: "cui7olrl0ubs73d8qpe0", Type: "subnet"}, {ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
}
|
||||
|
||||
func TestGetAccountSetings(t *testing.T) {
|
||||
ctx := context.TODO()
|
||||
|
||||
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
|
||||
assert.NoError(t, err)
|
||||
// err = loadSQL(ctx, s.pool, initDb)
|
||||
//assert.NoError(t, err)
|
||||
|
||||
settings, err := s.GetAccountSettings(ctx, "d5n27dafadhs73bt5ovg") //"ckd7ee2fic3c73dtendg")
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, nmdata.AccountSettingsInfo{
|
||||
PeerLoginExpirationEnabled: true,
|
||||
PeerLoginExpiration: 86400000000000 * time.Nanosecond,
|
||||
PeerInactivityExpirationEnabled: true,
|
||||
PeerInactivityExpiration: 600000000000 * time.Nanosecond,
|
||||
}, settings)
|
||||
}
|
||||
|
||||
func TestGetPostureChecks(t *testing.T) {
|
||||
ctx := context.TODO()
|
||||
|
||||
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
|
||||
assert.NoError(t, err)
|
||||
// err = loadSQL(ctx, s.pool, initDb)
|
||||
//assert.NoError(t, err)
|
||||
|
||||
checks, err := s.GetPostureChecks(ctx, "cdfcks2t2r9s73a58us0") //"ckd7ee2fic3c73dtendg")
|
||||
assert.NoError(t, err)
|
||||
|
||||
fmt.Print(checks)
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "test-group-1", PublicID: "public-id-1", Resources: []nmdata.Resource{{ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
// assert.Contains(t,
|
||||
// groups,
|
||||
// nmdata.Group{Name: "All", PublicID: "d9aejspvcsu517nkh4a0", Resources: []nmdata.Resource{{ID: "cui7olrl0ubs73d8qpe0", Type: "subnet"}, {ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
|
||||
// )
|
||||
}
|
||||
|
||||
func TestGetAllowedUserIds(t *testing.T) {
|
||||
ctx := context.TODO()
|
||||
|
||||
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
|
||||
assert.NoError(t, err)
|
||||
// err = loadSQL(ctx, s.pool, initDb)
|
||||
//assert.NoError(t, err)
|
||||
|
||||
userIds, groupToUserIds, err := s.GetAllowedUsers(ctx, "cus73sbl0ubs73cfoo90") //"ckd7ee2fic3c73dtendg")
|
||||
assert.NoError(t, err)
|
||||
assert.NotEmpty(t, userIds)
|
||||
assert.NotEmpty(t, groupToUserIds)
|
||||
}
|
||||
|
||||
func TestGetDnsSettings(t *testing.T) {
|
||||
ctx := context.TODO()
|
||||
|
||||
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
|
||||
assert.NoError(t, err)
|
||||
// err = loadSQL(ctx, s.pool, initDb)
|
||||
//assert.NoError(t, err)
|
||||
|
||||
set, err := s.GetDnsSettings(ctx, "ckvdmrqfic3c739ihh5g") //"ckd7ee2fic3c73dtendg")
|
||||
assert.NoError(t, err)
|
||||
assert.NotEmpty(t, set.DisabledManagementGroups)
|
||||
}
|
||||
65
management/internals/network_map_db/pgsql/nameserver.go
Normal file
65
management/internals/network_map_db/pgsql/nameserver.go
Normal file
@@ -0,0 +1,65 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"reflect"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
const (
|
||||
GetNameserversQuery = `
|
||||
select id, public_id, name, description, name_servers, groups, "primary", domains, enabled, search_domains_enabled
|
||||
from name_server_groups
|
||||
where account_id=$1
|
||||
`
|
||||
)
|
||||
|
||||
func (pg *PgStore) GetNameServerGroups(ctx context.Context, accountId string) ([]nmdata.NameServerGroup, error) {
|
||||
c, err := pg.Pool.Acquire(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetNameServerGroupsViaPgxConnection(ctx, c.Conn(), accountId)
|
||||
}
|
||||
|
||||
func GetNameServerGroupsViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.NameServerGroup, error) {
|
||||
rows, err := con.Query(ctx, GetNameserversQuery, accountId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
nsgroups, err := pgx.CollectRows(rows, pgx.RowToStructByName[nameserverGroup])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
toret := make([]nmdata.NameServerGroup, 0, len(nsgroups))
|
||||
for _, nsg := range nsgroups {
|
||||
group := nmdata.NameServerGroup{}
|
||||
err := networkmapdb.FromSqlTypesToSharedTypes(
|
||||
reflect.ValueOf(&nsg), reflect.ValueOf(&group))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
toret = append(toret, group)
|
||||
}
|
||||
return toret, nil
|
||||
}
|
||||
|
||||
type nameserverGroup struct {
|
||||
ID string
|
||||
PublicID sql.NullString
|
||||
Name sql.NullString
|
||||
Description sql.NullString
|
||||
NameServers json.RawMessage
|
||||
Groups json.RawMessage
|
||||
Primary sql.NullBool
|
||||
Domains json.RawMessage
|
||||
Enabled sql.NullBool
|
||||
SearchDomainsEnabled sql.NullBool
|
||||
}
|
||||
57
management/internals/network_map_db/pgsql/network.go
Normal file
57
management/internals/network_map_db/pgsql/network.go
Normal file
@@ -0,0 +1,57 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"reflect"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
const (
|
||||
GetNetworkQuery = `
|
||||
select network_identifier as identifier, network_net as net, network_net_v6 as net_v6, network_dns as dns, network_serial as serial
|
||||
from accounts
|
||||
where id=$1
|
||||
`
|
||||
)
|
||||
|
||||
func (pg *PgStore) GetNetwork(ctx context.Context, accountId string) (nmdata.Network, error) {
|
||||
c, err := pg.Pool.Acquire(ctx)
|
||||
if err != nil {
|
||||
return nmdata.Network{}, err
|
||||
}
|
||||
return GetNetworkViaPgxConnection(ctx, c.Conn(), accountId)
|
||||
}
|
||||
|
||||
func GetNetworkViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (nmdata.Network, error) {
|
||||
rows, err := con.Query(ctx, GetNetworkQuery, accountId)
|
||||
if err != nil {
|
||||
return nmdata.Network{}, err
|
||||
}
|
||||
|
||||
n, err := pgx.CollectOneRow(rows, pgx.RowToStructByName[accountnetwork])
|
||||
if err != nil {
|
||||
return nmdata.Network{}, err
|
||||
}
|
||||
|
||||
toret := nmdata.Network{}
|
||||
err = networkmapdb.FromSqlTypesToSharedTypes(
|
||||
reflect.ValueOf(&n), reflect.ValueOf(&toret))
|
||||
if err != nil {
|
||||
return nmdata.Network{}, err
|
||||
}
|
||||
|
||||
return toret, nil
|
||||
}
|
||||
|
||||
type accountnetwork struct {
|
||||
Identifier sql.NullString
|
||||
Net json.RawMessage
|
||||
NetV6 json.RawMessage
|
||||
Dns sql.NullString
|
||||
Serial sql.NullInt64
|
||||
}
|
||||
147
management/internals/network_map_db/pgsql/network_map_data.go
Normal file
147
management/internals/network_map_db/pgsql/network_map_data.go
Normal file
@@ -0,0 +1,147 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
func (pg *PgStore) GetNetworkMapData(ctx context.Context, accountId string) (*networkmap.NetworkMapData, error) {
|
||||
tx, err := pg.Pool.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
acctSettings, err := GetAccountSettingsViaPgxConnection(ctx, tx.Conn(), accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
dnsZones, err := GetAppliedZoneCandidatesViaPgxConnection(ctx, tx.Conn(), accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
groups, resourceToGroupIdx, err := GetGroupsViaPgxConnection(ctx, tx.Conn(), accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
nsGroups, err := GetNameServerGroupsViaPgxConnection(ctx, tx.Conn(), accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
networkResources, err := GetNetworkResourcesViaPgxConnection(ctx, tx.Conn(), accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
routers, err := GetNetworkRoutersViaPgxConnection(ctx, tx.Conn(), accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
network, err := GetNetworkViaPgxConnection(ctx, tx.Conn(), accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
peers, _, err := GetPeersViaPgxConnection(ctx, tx.Conn(), accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
policies, policyToDestinationResourceIdx, policyToDestinationGroupIdx, err := GetPoliciesViaPgxConnection(ctx, tx.Conn(), accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
postureChecks, err := GetPostureChecksViaPgxConnection(ctx, tx.Conn(), accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
routes, err := GetRoutesViaPgxConnection(ctx, tx.Conn(), accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
networkXIDToPublicID, err := GetNetworkXIDToPublicIdMapViaPgxConnection(ctx, tx.Conn(), accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
allowedUserIds, groupsToUserIds, err := GetAllowedUsersViaPgxConnection(ctx, tx.Conn(), accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
dnsSettings, err := GetDnsSettingsViaPgxConnection(ctx, tx.Conn(), accountId)
|
||||
if err != nil {
|
||||
return rollbackAndReturnError(ctx, tx, err)
|
||||
}
|
||||
|
||||
resourcePolicies := make(map[string][]*nmdata.Policy)
|
||||
for _, resource := range networkResources {
|
||||
if !resource.Enabled {
|
||||
continue
|
||||
}
|
||||
networkResourceGroups := resourceToGroupIdx[resource.ID]
|
||||
for _, policy := range policies {
|
||||
if !policy.Enabled {
|
||||
continue
|
||||
}
|
||||
if _, ok := policyToDestinationResourceIdx[policy.ID][resource.ID]; ok {
|
||||
resourcePolicies[resource.ID] = append(resourcePolicies[resource.ID], &policy) // TODO (dmitri) maybe use public id?
|
||||
break
|
||||
}
|
||||
if groupIds, ok := policyToDestinationGroupIdx[policy.ID]; ok {
|
||||
for networkResourceGroup := range networkResourceGroups {
|
||||
if _, ok := groupIds[networkResourceGroup]; ok {
|
||||
resourcePolicies[resource.ID] = append(resourcePolicies[resource.ID], &policy)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
err = tx.Commit(ctx)
|
||||
if err != nil {
|
||||
// TODO log and ignore?
|
||||
}
|
||||
|
||||
toret := networkmap.NetworkMapData{
|
||||
AccountSettings: &acctSettings,
|
||||
DNSSettings: &dnsSettings,
|
||||
Network: &network,
|
||||
Peers: toMap(peers, func(p nmdata.Peer) string { return p.ID }),
|
||||
Groups: toMap(groups, func(g nmdata.Group) string { return g.PublicID }),
|
||||
Policies: toSliceOfPtrs(policies),
|
||||
ResourcePolicies: resourcePolicies,
|
||||
Routes: toSliceOfPtrs(routes),
|
||||
Routers: routers,
|
||||
NameServerGroups: toSliceOfPtrs(nsGroups),
|
||||
NetworkResources: toSliceOfPtrs(networkResources),
|
||||
PostureChecks: toMap(postureChecks, func(pc nmdata.PostureChecks) string { return pc.ID }),
|
||||
AllowedUserIDs: allowedUserIds,
|
||||
GroupIDToUserIDs: groupsToUserIds,
|
||||
NetworkXIDToPublicID: networkXIDToPublicID, // TODO (dmitri) maybe we can switch to public ids everywhere?
|
||||
AppliedZoneCandidates: dnsZones,
|
||||
}
|
||||
|
||||
return &toret, nil
|
||||
}
|
||||
|
||||
func rollbackAndReturnError(ctx context.Context, tx pgx.Tx, err error) (*networkmap.NetworkMapData, error) {
|
||||
if errr := tx.Rollback(ctx); errr != nil {
|
||||
// TODO log and ignore?
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
func toMap[T any](all []T, id func(t T) string) map[string]*T {
|
||||
toret := make(map[string]*T, len(all))
|
||||
for _, t := range all {
|
||||
toret[id(t)] = &t
|
||||
}
|
||||
return toret
|
||||
}
|
||||
|
||||
func toSliceOfPtrs[T any](all []T) []*T {
|
||||
toret := make([]*T, len(all))
|
||||
for _, t := range all {
|
||||
toret = append(toret, &t)
|
||||
}
|
||||
return toret
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"reflect"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
const (
|
||||
GetNetworkResourcesQuery = `
|
||||
select id, network_id, account_id, public_id, name, description, type, domain, prefix, enabled
|
||||
from network_resources
|
||||
where account_id=$1
|
||||
`
|
||||
)
|
||||
|
||||
func (pg *PgStore) GetNetworkResources(ctx context.Context, accountId string) ([]nmdata.NetworkResource, error) {
|
||||
c, err := pg.Pool.Acquire(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetNetworkResourcesViaPgxConnection(ctx, c.Conn(), accountId)
|
||||
}
|
||||
|
||||
func GetNetworkResourcesViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.NetworkResource, error) {
|
||||
rows, err := con.Query(ctx, GetNetworkResourcesQuery, accountId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
netresorces, err := pgx.CollectRows(rows, pgx.RowToStructByName[networkresource])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
toret := make([]nmdata.NetworkResource, 0, len(netresorces))
|
||||
for _, nres := range netresorces {
|
||||
resource := nmdata.NetworkResource{}
|
||||
err := networkmapdb.FromSqlTypesToSharedTypes(
|
||||
reflect.ValueOf(&nres), reflect.ValueOf(&resource))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
toret = append(toret, resource)
|
||||
}
|
||||
return toret, nil
|
||||
}
|
||||
|
||||
type networkresource struct {
|
||||
ID string
|
||||
NetworkID sql.NullString
|
||||
AccountID sql.NullString
|
||||
PublicID sql.NullString
|
||||
Name sql.NullString
|
||||
Description sql.NullString
|
||||
Type sql.NullString
|
||||
Domain sql.NullString
|
||||
Prefix json.RawMessage
|
||||
Enabled sql.NullBool
|
||||
}
|
||||
85
management/internals/network_map_db/pgsql/network_router.go
Normal file
85
management/internals/network_map_db/pgsql/network_router.go
Normal file
@@ -0,0 +1,85 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"reflect"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
const (
|
||||
GetNetworkRouterQuery = `
|
||||
select public_id, peer, network_id, masquerade, metric, enabled,
|
||||
(
|
||||
select array_agg(group_peers.peer_id)
|
||||
from group_peers
|
||||
where group_peers.group_id in (select json_array_elements_text(peer_groups::json))
|
||||
) as peers_via_groups
|
||||
from network_routers
|
||||
where account_id=$1
|
||||
`
|
||||
)
|
||||
|
||||
func (pg *PgStore) GetNetworkRouters(ctx context.Context, accountId string) (map[string]map[string]*nmdata.NetworkRouter, error) {
|
||||
c, err := pg.Pool.Acquire(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetNetworkRoutersViaPgxConnection(ctx, c.Conn(), accountId)
|
||||
}
|
||||
|
||||
func GetNetworkRoutersViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (map[string]map[string]*nmdata.NetworkRouter, error) {
|
||||
rows, err := con.Query(ctx, GetNetworkRouterQuery, accountId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
routers, err := pgx.CollectRows(rows, pgx.RowToStructByName[networkrouter])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
toret := make(map[string]map[string]*nmdata.NetworkRouter)
|
||||
for _, router := range routers {
|
||||
if !router.Enabled.Bool {
|
||||
continue
|
||||
}
|
||||
|
||||
networkId := router.NetworkID.String
|
||||
if networkId == "" {
|
||||
return nil, fmt.Errorf("router with public_id %s doesn't have network_id set", router.PublicID.String)
|
||||
}
|
||||
|
||||
nmdatarouter := nmdata.NetworkRouter{}
|
||||
err := networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&router), reflect.ValueOf(&nmdatarouter))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if toret[networkId] == nil {
|
||||
toret[networkId] = make(map[string]*nmdata.NetworkRouter)
|
||||
}
|
||||
if router.Peer.String != "" {
|
||||
toret[networkId][router.Peer.String] = &nmdatarouter
|
||||
}
|
||||
for _, peerId := range router.PeersViaGroups {
|
||||
toret[networkId][peerId] = &nmdatarouter
|
||||
}
|
||||
}
|
||||
|
||||
return toret, nil
|
||||
}
|
||||
|
||||
type networkrouter struct {
|
||||
PublicID sql.NullString
|
||||
NetworkID sql.NullString `nmap:"skip"`
|
||||
Peer sql.NullString `nmap:"skip"`
|
||||
PeersViaGroups []string `nmap:"skip"`
|
||||
Masquerade sql.NullBool
|
||||
Metric sql.NullInt64
|
||||
Enabled sql.NullBool
|
||||
}
|
||||
52
management/internals/network_map_db/pgsql/networks.go
Normal file
52
management/internals/network_map_db/pgsql/networks.go
Normal file
@@ -0,0 +1,52 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
)
|
||||
|
||||
const (
|
||||
GetNetworksQuery = `
|
||||
select id, public_id
|
||||
from networks where account_id=$1
|
||||
`
|
||||
)
|
||||
|
||||
func (pg *PgStore) GetNetworks(ctx context.Context, accountId string) ([]network, error) {
|
||||
c, err := pg.Pool.Acquire(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetNetworksViaPgxConnection(ctx, c.Conn(), accountId)
|
||||
}
|
||||
|
||||
func GetNetworksViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]network, error) {
|
||||
rows, err := con.Query(ctx, GetGroupsQuery, accountId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return pgx.CollectRows(rows, pgx.RowToStructByName[network])
|
||||
}
|
||||
|
||||
func GetNetworkXIDToPublicIdMapViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (map[string]string, error) {
|
||||
networks, err := GetNetworksViaPgxConnection(ctx, con, accountId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
toret := make(map[string]string)
|
||||
for _, n := range networks {
|
||||
if n.PublicID.Valid {
|
||||
toret[n.ID] = n.PublicID.String
|
||||
}
|
||||
}
|
||||
|
||||
return toret, nil
|
||||
}
|
||||
|
||||
type network struct {
|
||||
ID string
|
||||
PublicID sql.NullString
|
||||
}
|
||||
146
management/internals/network_map_db/pgsql/peer.go
Normal file
146
management/internals/network_map_db/pgsql/peer.go
Normal file
@@ -0,0 +1,146 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"reflect"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
const (
|
||||
GetPeersQuery = `
|
||||
select id, key, ssh_key, dns_label, extra_dns_labels, user_id, ssh_enabled, login_expiration_enabled, last_login, ip, ipv6,
|
||||
peer_status_requires_approval, proxy_meta_embedded, proxy_meta_cluster,
|
||||
meta_wt_version, meta_go_os, meta_os_version, meta_kernel_version, meta_network_addresses, meta_files, meta_capabilities, meta_flags, meta_sync_message_version,
|
||||
location_country_code, location_city_name, location_connection_ip
|
||||
from peers
|
||||
where account_id = $1
|
||||
`
|
||||
)
|
||||
|
||||
func (pg *PgStore) GetPeers(ctx context.Context, accountId string) ([]nmdata.Peer, map[string]*nmdata.Peer, error) {
|
||||
c, err := pg.Pool.Acquire(ctx)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
return GetPeersViaPgxConnection(ctx, c.Conn(), accountId)
|
||||
}
|
||||
|
||||
func GetPeersViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.Peer, map[string]*nmdata.Peer, error) {
|
||||
rows, err := con.Query(ctx, GetPeersQuery, accountId)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
peers, err := pgx.CollectRows(rows, pgx.RowToStructByName[peer])
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
toret := make([]nmdata.Peer, 0, len(peers))
|
||||
clusterToPeerIdx := make(map[string]*nmdata.Peer)
|
||||
for _, p := range peers {
|
||||
dp := nmdata.Peer{}
|
||||
err := networkmapdb.FromSqlTypesToSharedTypes(
|
||||
reflect.ValueOf(&p), reflect.ValueOf(&dp))
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
if p.ProxyMetaEmbedded.Valid {
|
||||
dp.ProxyMeta.Embedded = p.ProxyMetaEmbedded.Bool
|
||||
}
|
||||
if dp.ProxyMeta.Embedded {
|
||||
clusterToPeerIdx[p.ProxyMetaCluster.String] = &dp
|
||||
}
|
||||
if p.MetaWtVersion.Valid {
|
||||
dp.Meta.WtVersion = p.MetaWtVersion.String
|
||||
}
|
||||
if p.MetaSyncMessageVersion.Valid {
|
||||
dp.Meta.SyncMessageVersion = int(p.MetaSyncMessageVersion.Int64)
|
||||
}
|
||||
if p.MetaGoOS.Valid {
|
||||
dp.Meta.GoOS = p.MetaGoOS.String
|
||||
}
|
||||
if p.MetaOSVersion.Valid {
|
||||
dp.Meta.OSVersion = p.MetaOSVersion.String
|
||||
}
|
||||
if p.MetaKernelVersion.Valid {
|
||||
dp.Meta.KernelVersion = p.MetaKernelVersion.String
|
||||
}
|
||||
if p.LocationCountryCode.Valid {
|
||||
dp.Location.CountryCode = p.LocationCountryCode.String
|
||||
}
|
||||
if p.LocationCityName.Valid {
|
||||
dp.Location.CityName = p.LocationCityName.String
|
||||
}
|
||||
if p.LocationConnectionIp != nil {
|
||||
err := json.Unmarshal(p.LocationConnectionIp, &dp.Location.ConnectionIP)
|
||||
if err != nil {
|
||||
return toret, nil, err
|
||||
}
|
||||
}
|
||||
if p.MetaFiles != nil {
|
||||
err := json.Unmarshal(p.MetaFiles, &dp.Meta.Files)
|
||||
if err != nil {
|
||||
return toret, nil, err
|
||||
}
|
||||
}
|
||||
if p.MetaCapabilities != nil {
|
||||
err := json.Unmarshal(p.MetaCapabilities, &dp.Meta.Capabilities)
|
||||
if err != nil {
|
||||
return toret, nil, err
|
||||
}
|
||||
}
|
||||
if p.MetaFlags != nil {
|
||||
err := json.Unmarshal(p.MetaFlags, &dp.Meta.Flags)
|
||||
if err != nil {
|
||||
return toret, nil, err
|
||||
}
|
||||
}
|
||||
if p.MetaNetworkAddresses != nil {
|
||||
err := json.Unmarshal(p.MetaNetworkAddresses, &dp.Meta.NetworkAddresses)
|
||||
if err != nil {
|
||||
return toret, nil, err
|
||||
}
|
||||
}
|
||||
|
||||
toret = append(toret, dp)
|
||||
}
|
||||
|
||||
return toret, clusterToPeerIdx, nil
|
||||
}
|
||||
|
||||
// TODO add support for creating struct fields from denormalized fields
|
||||
type peer struct {
|
||||
ID string
|
||||
Key sql.NullString
|
||||
SSHKey sql.NullString
|
||||
DNSLabel sql.NullString
|
||||
ExtraDNSLabels json.RawMessage
|
||||
UserID sql.NullString
|
||||
LastLogin sql.NullTime
|
||||
SSHEnabled sql.NullBool
|
||||
LoginExpirationEnabled sql.NullBool
|
||||
PeerStatusRequiresApproval sql.NullBool `nmap:"map_to:RequiresApproval"`
|
||||
ProxyMetaEmbedded sql.NullBool `nmap:"skip"`
|
||||
ProxyMetaCluster sql.NullString `nmap:"skip"`
|
||||
IP json.RawMessage
|
||||
IPv6 json.RawMessage
|
||||
LocationConnectionIp json.RawMessage `nmap:"skip"`
|
||||
MetaFiles json.RawMessage `nmap:"skip"`
|
||||
MetaCapabilities json.RawMessage `nmap:"skip"`
|
||||
MetaFlags json.RawMessage `nmap:"skip"`
|
||||
MetaNetworkAddresses json.RawMessage `nmap:"skip"`
|
||||
MetaWtVersion sql.NullString `nmap:"skip"`
|
||||
MetaGoOS sql.NullString `nmap:"skip"`
|
||||
MetaOSVersion sql.NullString `nmap:"skip"`
|
||||
MetaKernelVersion sql.NullString `nmap:"skip"`
|
||||
MetaSyncMessageVersion sql.NullInt64 `nmap:"skip"`
|
||||
LocationCountryCode sql.NullString `nmap:"skip"`
|
||||
LocationCityName sql.NullString `nmap:"skip"`
|
||||
}
|
||||
56
management/internals/network_map_db/pgsql/pg_store.go
Normal file
56
management/internals/network_map_db/pgsql/pg_store.go
Normal file
@@ -0,0 +1,56 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
)
|
||||
|
||||
const (
|
||||
pgMaxConnections = 30
|
||||
pgMinConnections = 1
|
||||
pgMaxConnLifetime = 60 * time.Minute
|
||||
pgHealthCheckPeriod = 1 * time.Minute
|
||||
)
|
||||
|
||||
var _ networkmapdb.NetworkMapDBStore = &PgStore{}
|
||||
|
||||
type PgStore struct {
|
||||
Pool *pgxpool.Pool
|
||||
}
|
||||
|
||||
func NewPostgresqlStore(ctx context.Context, dsn string) (*PgStore, error) {
|
||||
pool, err := connectToPgDb(context.Background(), dsn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &PgStore{Pool: pool}, nil
|
||||
}
|
||||
|
||||
func connectToPgDb(ctx context.Context, dsn string) (*pgxpool.Pool, error) {
|
||||
config, err := pgxpool.ParseConfig(dsn)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("unable to parse database config: %w", err)
|
||||
}
|
||||
|
||||
config.MaxConns = pgMaxConnections
|
||||
config.MinConns = pgMinConnections
|
||||
config.MaxConnLifetime = pgMaxConnLifetime
|
||||
config.HealthCheckPeriod = pgHealthCheckPeriod
|
||||
|
||||
pool, err := pgxpool.NewWithConfig(ctx, config)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("unable to create connection pool: %w", err)
|
||||
}
|
||||
|
||||
if err := pool.Ping(ctx); err != nil {
|
||||
pool.Close()
|
||||
return nil, fmt.Errorf("unable to ping database: %w", err)
|
||||
}
|
||||
|
||||
return pool, nil
|
||||
}
|
||||
164
management/internals/network_map_db/pgsql/policy.go
Normal file
164
management/internals/network_map_db/pgsql/policy.go
Normal file
@@ -0,0 +1,164 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"reflect"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
const (
|
||||
GetPoliciesQuery = `
|
||||
select p.id, p.public_id, p.enabled, p.source_posture_checks, pr.enabled as rule_enabled, pr.action, pr.protocol, pr.bidirectional,
|
||||
pr.sources, pr.destinations, pr.source_resource, pr.destination_resource, pr.ports, pr.port_ranges,
|
||||
pr.authorized_groups, pr.authorized_user
|
||||
from policies as p
|
||||
left join policy_rules as pr on p.id = pr.policy_id
|
||||
where account_id=$1
|
||||
`
|
||||
)
|
||||
|
||||
func (pg *PgStore) GetPolicies(ctx context.Context, accountId string) ([]nmdata.Policy, map[string]map[string]any, map[string]map[string]any, error) {
|
||||
c, err := pg.Pool.Acquire(ctx)
|
||||
if err != nil {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
return GetPoliciesViaPgxConnection(ctx, c.Conn(), accountId)
|
||||
}
|
||||
|
||||
func GetPoliciesViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.Policy, map[string]map[string]any, map[string]map[string]any, error) {
|
||||
rows, err := con.Query(ctx, GetPoliciesQuery, accountId)
|
||||
if err != nil {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
|
||||
policies, err := pgx.CollectRows(rows, pgx.RowToStructByName[policy])
|
||||
if err != nil {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
|
||||
toret := make([]nmdata.Policy, 0, len(policies))
|
||||
policyToDestinationResourceIdx := make(map[string]map[string]any) // policy id to destination resource id
|
||||
policyToDestinationGroupIdx := make(map[string]map[string]any) // policy id to destination group id
|
||||
for _, p := range policies {
|
||||
policy := nmdata.Policy{}
|
||||
err := networkmapdb.FromSqlTypesToSharedTypes(
|
||||
reflect.ValueOf(&p), reflect.ValueOf(&policy))
|
||||
if err != nil {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
|
||||
var policyRule *nmdata.PolicyRule
|
||||
pr := func() *nmdata.PolicyRule {
|
||||
if policyRule != nil {
|
||||
return policyRule
|
||||
}
|
||||
|
||||
policyRule = &nmdata.PolicyRule{}
|
||||
return policyRule
|
||||
}
|
||||
|
||||
if p.RuleEnabled.Valid {
|
||||
pr().Enabled = p.RuleEnabled.Bool
|
||||
}
|
||||
if p.Action.Valid {
|
||||
pr().Action = p.Action.String
|
||||
}
|
||||
if p.Protocol.Valid {
|
||||
pr().Protocol = p.Protocol.String
|
||||
}
|
||||
if p.Bidirectional.Valid {
|
||||
pr().Bidirectional = p.Bidirectional.Bool
|
||||
}
|
||||
if len(p.Sources) > 0 {
|
||||
err := json.Unmarshal([]byte(p.Sources), &pr().Sources)
|
||||
if err != nil {
|
||||
return toret, nil, nil, err
|
||||
}
|
||||
}
|
||||
if len(p.Destinations) > 0 {
|
||||
err := json.Unmarshal([]byte(p.Destinations), &pr().Destinations)
|
||||
if err != nil {
|
||||
return toret, nil, nil, err
|
||||
}
|
||||
|
||||
for _, dst := range pr().Destinations {
|
||||
if _, ok := policyToDestinationGroupIdx[p.ID]; !ok {
|
||||
policyToDestinationGroupIdx[p.ID] = make(map[string]any)
|
||||
}
|
||||
policyToDestinationGroupIdx[p.ID][dst] = struct{}{}
|
||||
}
|
||||
}
|
||||
if len(p.SourceResource) > 0 {
|
||||
err := json.Unmarshal([]byte(p.SourceResource), &pr().SourceResource)
|
||||
if err != nil {
|
||||
return toret, nil, nil, err
|
||||
}
|
||||
}
|
||||
if len(p.DestinationResource) > 0 {
|
||||
err := json.Unmarshal([]byte(p.DestinationResource), &pr().DestinationResource)
|
||||
if err != nil {
|
||||
return toret, nil, nil, err
|
||||
}
|
||||
|
||||
if _, ok := policyToDestinationResourceIdx[p.ID]; !ok {
|
||||
policyToDestinationResourceIdx[p.ID] = make(map[string]any)
|
||||
}
|
||||
policyToDestinationResourceIdx[p.ID][pr().DestinationResource.ID] = struct{}{}
|
||||
}
|
||||
if len(p.Ports) > 0 {
|
||||
err := json.Unmarshal([]byte(p.Ports), &pr().Ports)
|
||||
if err != nil {
|
||||
return toret, nil, nil, err
|
||||
}
|
||||
}
|
||||
if len(p.PortRanges) > 0 {
|
||||
err := json.Unmarshal([]byte(p.PortRanges), &pr().PortRanges)
|
||||
if err != nil {
|
||||
return toret, nil, nil, err
|
||||
}
|
||||
}
|
||||
if len(p.AuthorizedGroups) > 0 {
|
||||
err := json.Unmarshal([]byte(p.AuthorizedGroups), &pr().AuthorizedGroups)
|
||||
if err != nil {
|
||||
return toret, nil, nil, err
|
||||
}
|
||||
}
|
||||
if p.AuthorizedUser.Valid {
|
||||
pr().AuthorizedUser = p.AuthorizedUser.String
|
||||
}
|
||||
|
||||
if policyRule != nil {
|
||||
policyRule.ID = p.ID
|
||||
policyRule.PolicyID = p.ID
|
||||
policy.Rules = []*nmdata.PolicyRule{policyRule}
|
||||
}
|
||||
|
||||
toret = append(toret, policy)
|
||||
}
|
||||
|
||||
return toret, policyToDestinationResourceIdx, policyToDestinationGroupIdx, err
|
||||
}
|
||||
|
||||
type policy struct {
|
||||
ID string
|
||||
PublicID sql.NullString
|
||||
SourcePostureChecks json.RawMessage
|
||||
Enabled sql.NullBool
|
||||
RuleEnabled sql.NullBool `nmap:"skip"`
|
||||
Bidirectional sql.NullBool `nmap:"skip"`
|
||||
Action sql.NullString `nmap:"skip"`
|
||||
Protocol sql.NullString `nmap:"skip"`
|
||||
Sources json.RawMessage `nmap:"skip"`
|
||||
Destinations json.RawMessage `nmap:"skip"`
|
||||
SourceResource json.RawMessage `nmap:"skip"`
|
||||
DestinationResource json.RawMessage `nmap:"skip"`
|
||||
Ports json.RawMessage `nmap:"skip"`
|
||||
PortRanges json.RawMessage `nmap:"skip"`
|
||||
AuthorizedGroups json.RawMessage `nmap:"skip"`
|
||||
AuthorizedUser sql.NullString `nmap:"skip"`
|
||||
}
|
||||
56
management/internals/network_map_db/pgsql/posture.go
Normal file
56
management/internals/network_map_db/pgsql/posture.go
Normal file
@@ -0,0 +1,56 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"reflect"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
const (
|
||||
GetPostureChecksQuery = `
|
||||
select public_id as id, checks
|
||||
from posture_checks
|
||||
where account_id=$1
|
||||
`
|
||||
)
|
||||
|
||||
func (pg *PgStore) GetPostureChecks(ctx context.Context, accountId string) ([]nmdata.PostureChecks, error) {
|
||||
c, err := pg.Pool.Acquire(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetPostureChecksViaPgxConnection(ctx, c.Conn(), accountId)
|
||||
}
|
||||
|
||||
func GetPostureChecksViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.PostureChecks, error) {
|
||||
rows, err := con.Query(ctx, GetPostureChecksQuery, accountId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
checks, err := pgx.CollectRows(rows, pgx.RowToStructByName[posturechecks])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
toret := make([]nmdata.PostureChecks, 0, len(checks))
|
||||
for _, c := range checks {
|
||||
checks := nmdata.PostureChecks{}
|
||||
err := networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&c), reflect.ValueOf(&checks))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
toret = append(toret, checks)
|
||||
}
|
||||
|
||||
return toret, nil
|
||||
}
|
||||
|
||||
type posturechecks struct {
|
||||
ID string
|
||||
Checks json.RawMessage
|
||||
}
|
||||
75
management/internals/network_map_db/pgsql/route.go
Normal file
75
management/internals/network_map_db/pgsql/route.go
Normal file
@@ -0,0 +1,75 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"reflect"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
const (
|
||||
GetRoutesQuery = `
|
||||
select id, account_id, public_id, network, domains, keep_route, net_id, description,
|
||||
peer, peer as peer_id, peer_groups, network_type, masquerade, metric, enabled,
|
||||
groups, access_control_groups, skip_auto_apply
|
||||
from routes
|
||||
where account_id=$1
|
||||
`
|
||||
)
|
||||
|
||||
func (pg *PgStore) GetRoutes(ctx context.Context, accountId string) ([]nmdata.Route, error) {
|
||||
c, err := pg.Pool.Acquire(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetRoutesViaPgxConnection(ctx, c.Conn(), accountId)
|
||||
}
|
||||
|
||||
func GetRoutesViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.Route, error) {
|
||||
rows, err := con.Query(ctx, GetRoutesQuery, accountId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
routes, err := pgx.CollectRows(rows, pgx.RowToStructByName[route])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
toret := make([]nmdata.Route, 0, len(routes))
|
||||
for _, r := range routes {
|
||||
route := nmdata.Route{}
|
||||
err := networkmapdb.FromSqlTypesToSharedTypes(
|
||||
reflect.ValueOf(&r), reflect.ValueOf(&route))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
toret = append(toret, route)
|
||||
}
|
||||
return toret, nil
|
||||
}
|
||||
|
||||
type route struct {
|
||||
ID string
|
||||
AccountID sql.NullString
|
||||
PublicID sql.NullString
|
||||
Network json.RawMessage
|
||||
Domains json.RawMessage
|
||||
KeepRoute sql.NullBool
|
||||
NetID sql.NullString
|
||||
Description sql.NullString
|
||||
Peer sql.NullString
|
||||
PeerID sql.NullString
|
||||
PeerGroups json.RawMessage
|
||||
NetworkType sql.NullInt64
|
||||
Masquerade sql.NullBool
|
||||
Metric sql.NullInt64
|
||||
Enabled sql.NullBool
|
||||
Groups json.RawMessage
|
||||
AccessControlGroups json.RawMessage
|
||||
SkipAutoApply sql.NullBool
|
||||
}
|
||||
@@ -0,0 +1,219 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"reflect"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestNullStringSupport(t *testing.T) {
|
||||
src := withNullString{Name: sql.NullString{String: "string", Valid: true}}
|
||||
dst := withString{}
|
||||
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
|
||||
assert.Equal(t, withString{Name: "string"}, dst)
|
||||
|
||||
src = withNullString{Name: sql.NullString{Valid: false}}
|
||||
dst = withString{}
|
||||
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
|
||||
assert.Equal(t, withString{Name: ""}, dst)
|
||||
}
|
||||
|
||||
func TestNullBoolSupport(t *testing.T) {
|
||||
src := withNullBool{TrueOrFalse: sql.NullBool{Bool: true, Valid: true}}
|
||||
dst := withBool{}
|
||||
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
|
||||
assert.Equal(t, withBool{TrueOrFalse: true}, dst)
|
||||
|
||||
}
|
||||
|
||||
func TestRawJsonSupport(t *testing.T) {
|
||||
jb, _ := json.Marshal(embeddedS{Name: "blob-name", SomeField: 1})
|
||||
src := withRawJson{Blob: json.RawMessage(jb)}
|
||||
dst := fromJson{}
|
||||
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
|
||||
assert.Equal(t, fromJson{Blob: embeddedS{Name: "blob-name", SomeField: 1}}, dst)
|
||||
|
||||
src1 := withRawJson{}
|
||||
dst1 := fromJson{}
|
||||
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src1), reflect.ValueOf(&dst1)))
|
||||
assert.Equal(t, fromJson{}, dst1)
|
||||
}
|
||||
|
||||
func TestShouldSkipTag(t *testing.T) {
|
||||
src5 := withSkipTag{Field: "shouldskip"}
|
||||
dst5 := emptySkipTagTarget{}
|
||||
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src5), reflect.ValueOf(&dst5)))
|
||||
assert.Equal(t, emptySkipTagTarget{}, dst5)
|
||||
|
||||
}
|
||||
|
||||
func TestMapToTag(t *testing.T) {
|
||||
src6 := withMapToTag{Field: "fieldvalue"}
|
||||
dst6 := mapToTagTarget{}
|
||||
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src6), reflect.ValueOf(&dst6)))
|
||||
assert.Equal(t, mapToTagTarget{AnotherField: "fieldvalue"}, dst6)
|
||||
}
|
||||
|
||||
func TestNullableInt64Support(t *testing.T) {
|
||||
src := withInt64{Field: sql.NullInt64{Int64: int64(1), Valid: true}}
|
||||
dst := int64Target{}
|
||||
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
|
||||
assert.Equal(t, int64Target{Field: 1}, dst)
|
||||
}
|
||||
|
||||
func TestNullableTimeSupport(t *testing.T) {
|
||||
now := time.Now()
|
||||
src := withNullableTime{Field: sql.NullTime{Time: now, Valid: true}}
|
||||
dst := nullableTimeTarget{}
|
||||
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
|
||||
assert.Equal(t, nullableTimeTarget{Field: now}, dst)
|
||||
}
|
||||
|
||||
func TestNullableTimePointerSupport(t *testing.T) {
|
||||
now := time.Now()
|
||||
src := withNullableTime{Field: sql.NullTime{Time: now, Valid: true}}
|
||||
dst := nullableTimePointerTarget{}
|
||||
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
|
||||
assert.Equal(t, nullableTimePointerTarget{Field: &now}, dst)
|
||||
}
|
||||
|
||||
func TestStringSLiceSupport(t *testing.T) {
|
||||
src := withStringSlice{Field: []string{"one"}}
|
||||
dst := withStringSlice{}
|
||||
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
|
||||
assert.Equal(t, withStringSlice{Field: []string{"one"}}, dst)
|
||||
}
|
||||
|
||||
func TestNullStringSLiceSupport(t *testing.T) {
|
||||
src := withStringSlice{}
|
||||
dst := withStringSlice{}
|
||||
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
|
||||
assert.Equal(t, withStringSlice{}, dst)
|
||||
}
|
||||
|
||||
func TestWithMultipleFields(t *testing.T) {
|
||||
now := time.Now()
|
||||
src := withMultipleFields{
|
||||
Field1: sql.NullString{String: "aaa", Valid: true},
|
||||
Field2: sql.NullBool{Bool: true, Valid: true},
|
||||
Field3: sql.NullTime{Time: now, Valid: true},
|
||||
Field4: sql.NullInt64{Int64: 1, Valid: true},
|
||||
Field5: "another",
|
||||
}
|
||||
dst := multipleFieldsTarget{}
|
||||
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
|
||||
assert.Equal(t, multipleFieldsTarget{
|
||||
Field1: "aaa",
|
||||
Field2: true,
|
||||
Field3: now,
|
||||
Field4: 1,
|
||||
Field5: "another",
|
||||
}, dst)
|
||||
}
|
||||
|
||||
func TestEmptyPublicIdsFilled(t *testing.T) {
|
||||
src := withEmptyPublicIds{}
|
||||
dst := emptyPublicIdTarget{}
|
||||
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
|
||||
assert.NotEmpty(t, dst.PublicID)
|
||||
assert.NotEmpty(t, dst.PublicId)
|
||||
}
|
||||
|
||||
type withNullString struct {
|
||||
Name sql.NullString
|
||||
}
|
||||
|
||||
type withString struct {
|
||||
Name string
|
||||
}
|
||||
|
||||
type withMultipleFields struct {
|
||||
Field1 sql.NullString
|
||||
Field2 sql.NullBool
|
||||
Field3 sql.NullTime
|
||||
Field4 sql.NullInt64
|
||||
Field5 string
|
||||
}
|
||||
|
||||
type multipleFieldsTarget struct {
|
||||
Field1 string
|
||||
Field2 bool
|
||||
Field3 time.Time
|
||||
Field4 int64
|
||||
Field5 string
|
||||
}
|
||||
|
||||
type withNullBool struct {
|
||||
TrueOrFalse sql.NullBool
|
||||
}
|
||||
|
||||
type withBool struct {
|
||||
TrueOrFalse bool
|
||||
}
|
||||
|
||||
type withRawJson struct {
|
||||
Blob json.RawMessage
|
||||
}
|
||||
|
||||
type embeddedS struct {
|
||||
Name string
|
||||
SomeField int
|
||||
}
|
||||
type fromJson struct {
|
||||
Blob embeddedS
|
||||
}
|
||||
|
||||
type withSkipTag struct {
|
||||
Field string `nmap:"skip"`
|
||||
}
|
||||
|
||||
type emptySkipTagTarget struct {
|
||||
Field string
|
||||
}
|
||||
|
||||
type withMapToTag struct {
|
||||
Field string `nmap:"map_to:AnotherField"`
|
||||
}
|
||||
|
||||
type mapToTagTarget struct {
|
||||
AnotherField string
|
||||
}
|
||||
|
||||
type withInt64 struct {
|
||||
Field sql.NullInt64
|
||||
}
|
||||
|
||||
type int64Target struct {
|
||||
Field int
|
||||
}
|
||||
|
||||
type withNullableTime struct {
|
||||
Field sql.NullTime
|
||||
}
|
||||
|
||||
type nullableTimeTarget struct {
|
||||
Field time.Time
|
||||
}
|
||||
|
||||
type nullableTimePointerTarget struct {
|
||||
Field *time.Time
|
||||
}
|
||||
|
||||
type withStringSlice struct {
|
||||
Field []string
|
||||
}
|
||||
|
||||
type withEmptyPublicIds struct {
|
||||
PublicID sql.NullString
|
||||
PublicId sql.NullString
|
||||
}
|
||||
|
||||
type emptyPublicIdTarget struct {
|
||||
PublicID string
|
||||
PublicId string
|
||||
}
|
||||
57
management/internals/network_map_db/pgsql/user.go
Normal file
57
management/internals/network_map_db/pgsql/user.go
Normal file
@@ -0,0 +1,57 @@
|
||||
package networkmap_pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
)
|
||||
|
||||
const (
|
||||
GetAllowedUserIdsQuery = `
|
||||
select id, auto_groups
|
||||
from users
|
||||
where account_id=$1 and not blocked and not is_service_user
|
||||
`
|
||||
)
|
||||
|
||||
func (pg *PgStore) GetAllowedUsers(ctx context.Context, accountId string) (map[string]struct{}, map[string][]string, error) {
|
||||
c, err := pg.Pool.Acquire(ctx)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
return GetAllowedUsersViaPgxConnection(ctx, c.Conn(), accountId)
|
||||
}
|
||||
|
||||
func GetAllowedUsersViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (map[string]struct{}, map[string][]string, error) {
|
||||
rows, err := con.Query(ctx, GetAllowedUserIdsQuery, accountId)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
users, err := pgx.CollectRows(rows, pgx.RowToStructByName[user])
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
userIdIdx := make(map[string]struct{})
|
||||
groupIdToUserIds := make(map[string][]string)
|
||||
for _, user := range users {
|
||||
userIdIdx[user.ID] = struct{}{}
|
||||
|
||||
var groupIds []string
|
||||
if err := json.Unmarshal(user.AutoGroups, &groupIds); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
for _, groupId := range groupIds {
|
||||
groupIdToUserIds[groupId] = append(groupIdToUserIds[groupId], user.ID)
|
||||
}
|
||||
}
|
||||
|
||||
return userIdIdx, groupIdToUserIds, nil
|
||||
}
|
||||
|
||||
type user struct {
|
||||
ID string
|
||||
AutoGroups json.RawMessage
|
||||
}
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"crypto/tls"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"os"
|
||||
"slices"
|
||||
"time"
|
||||
|
||||
@@ -24,13 +25,15 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/encryption"
|
||||
"github.com/netbirdio/netbird/formatter/hook"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
|
||||
accesslogsmanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs/manager"
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
|
||||
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
activitystore "github.com/netbirdio/netbird/management/server/activity/store"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
nbcache "github.com/netbirdio/netbird/management/server/cache"
|
||||
nbContext "github.com/netbirdio/netbird/management/server/context"
|
||||
nbhttp "github.com/netbirdio/netbird/management/server/http"
|
||||
@@ -99,6 +102,22 @@ func (s *BaseServer) Store() store.Store {
|
||||
})
|
||||
}
|
||||
|
||||
func (s *BaseServer) NetworkMapStore() networkmapdb.NetworkMapDBStore {
|
||||
return Create(s, func() networkmapdb.NetworkMapDBStore {
|
||||
dsn := os.Getenv("NETBIRD_NMAP_STORE_DSN") // Todo: this needs to be hoocked up properly
|
||||
if dsn == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
store, err := networkmap_pgsql.NewPostgresqlStore(context.Background(), dsn)
|
||||
if err != nil {
|
||||
log.Fatalf("failed to create network map store: %v", err)
|
||||
}
|
||||
|
||||
return store
|
||||
})
|
||||
}
|
||||
|
||||
func (s *BaseServer) EventStore() activity.Store {
|
||||
return Create(s, func() activity.Store {
|
||||
var err error
|
||||
|
||||
@@ -123,7 +123,7 @@ func (s *BaseServer) EphemeralManager() ephemeral.Manager {
|
||||
|
||||
func (s *BaseServer) NetworkMapController() network_map.Controller {
|
||||
return Create(s, func() network_map.Controller {
|
||||
return nmapcontroller.NewController(context.Background(), s.Store(), s.Metrics(), s.PeersUpdateManager(), s.AccountRequestBuffer(), s.IntegratedValidator(), s.SettingsManager(), s.DNSDomain(), s.ProxyController(), s.EphemeralManager(), s.Config)
|
||||
return nmapcontroller.NewController(context.Background(), s.Store(), s.Metrics(), s.PeersUpdateManager(), s.AccountRequestBuffer(), s.IntegratedValidator(), s.SettingsManager(), s.DNSDomain(), s.ProxyController(), s.EphemeralManager(), s.Config, s.NetworkMapStore())
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -4,10 +4,9 @@ import (
|
||||
"encoding/base64"
|
||||
"strconv"
|
||||
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
nbroute "github.com/netbirdio/netbird/route"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap"
|
||||
nmdata "github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
@@ -82,6 +81,7 @@ func EncodeNetworkMapEnvelope(in ComponentsEnvelopeInput) *proto.NetworkMapEnvel
|
||||
enc := newComponentEncoder(c)
|
||||
enc.indexAllPeers()
|
||||
routerIdxs := enc.indexRouterPeers(c.RouterPeers)
|
||||
enc.indexAllNetworkResources()
|
||||
|
||||
// Phase 2: gather every policy that any consumer references (peer-pair
|
||||
// policies + resource-only policies) so encodeResourcePoliciesMap can
|
||||
@@ -103,7 +103,6 @@ func EncodeNetworkMapEnvelope(in ComponentsEnvelopeInput) *proto.NetworkMapEnvel
|
||||
DnsSettings: enc.encodeDNSSettings(c.DNSSettings),
|
||||
DnsDomain: in.DNSDomain,
|
||||
CustomZoneDomain: c.CustomZoneDomain,
|
||||
AgentVersions: enc.agentVersions,
|
||||
Peers: enc.peers,
|
||||
RouterPeerIndexes: routerIdxs,
|
||||
Policies: policies,
|
||||
@@ -128,7 +127,7 @@ func EncodeNetworkMapEnvelope(in ComponentsEnvelopeInput) *proto.NetworkMapEnvel
|
||||
// networkSerial returns c.Network.CurrentSerial() with a nil guard. The
|
||||
// production path always populates c.Network, but the encoder is exported
|
||||
// and a hand-built components struct may omit it.
|
||||
func networkSerial(n *types.Network) uint64 {
|
||||
func networkSerial(n *nmdata.Network) uint64 {
|
||||
if n == nil {
|
||||
return 0
|
||||
}
|
||||
@@ -141,16 +140,15 @@ type componentEncoder struct {
|
||||
peerOrder map[string]uint32
|
||||
peers []*proto.PeerCompact
|
||||
|
||||
agentVersionOrder map[string]uint32
|
||||
agentVersions []string
|
||||
networkIdToPublicId map[string]string
|
||||
}
|
||||
|
||||
func newComponentEncoder(c *types.NetworkMapComponents) *componentEncoder {
|
||||
return &componentEncoder{
|
||||
components: c,
|
||||
peerOrder: make(map[string]uint32, len(c.Peers)),
|
||||
peers: make([]*proto.PeerCompact, 0, len(c.Peers)),
|
||||
agentVersionOrder: make(map[string]uint32),
|
||||
components: c,
|
||||
peerOrder: make(map[string]uint32, len(c.Peers)),
|
||||
peers: make([]*proto.PeerCompact, 0, len(c.Peers)),
|
||||
networkIdToPublicId: make(map[string]string),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -163,7 +161,7 @@ func (e *componentEncoder) indexAllPeers() {
|
||||
}
|
||||
}
|
||||
|
||||
func (e *componentEncoder) appendPeer(p *types.ComponentPeer) uint32 {
|
||||
func (e *componentEncoder) appendPeer(p *nmdata.Peer) uint32 {
|
||||
if idx, ok := e.peerOrder[p.ID]; ok {
|
||||
return idx
|
||||
}
|
||||
@@ -177,7 +175,7 @@ func (e *componentEncoder) appendPeer(p *types.ComponentPeer) uint32 {
|
||||
// (c.RouterPeers may contain peers not in c.Peers when validation rules drop
|
||||
// them) and returns their wire indexes for the RouterPeerIndexes field. Must
|
||||
// run before any encoder that resolves peer ids via e.peerOrder.
|
||||
func (e *componentEncoder) indexRouterPeers(routers map[string]*types.ComponentPeer) []uint32 {
|
||||
func (e *componentEncoder) indexRouterPeers(routers map[string]*nmdata.Peer) []uint32 {
|
||||
if len(routers) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -191,6 +189,15 @@ func (e *componentEncoder) indexRouterPeers(routers map[string]*types.ComponentP
|
||||
return out
|
||||
}
|
||||
|
||||
func (e *componentEncoder) indexAllNetworkResources() {
|
||||
for _, r := range e.components.NetworkResources {
|
||||
if !r.Enabled {
|
||||
continue
|
||||
}
|
||||
e.networkIdToPublicId[r.ID] = r.PublicID
|
||||
}
|
||||
}
|
||||
|
||||
func (e *componentEncoder) encodeGroups() []*proto.GroupCompact {
|
||||
if len(e.components.Groups) == 0 {
|
||||
return nil
|
||||
@@ -204,10 +211,20 @@ func (e *componentEncoder) encodeGroups() []*proto.GroupCompact {
|
||||
peerIdxs = append(peerIdxs, idx)
|
||||
}
|
||||
}
|
||||
|
||||
groupCompactResources := func() []*proto.ResourceCompact {
|
||||
var toret []*proto.ResourceCompact
|
||||
for _, r := range g.Resources {
|
||||
toret = append(toret, e.resourceToProto(r))
|
||||
}
|
||||
return toret
|
||||
}
|
||||
|
||||
out = append(out, &proto.GroupCompact{
|
||||
Id: g.PublicID,
|
||||
PeerIndexes: peerIdxs,
|
||||
IsAll: g.IsGroupAll(),
|
||||
Resources: groupCompactResources(),
|
||||
})
|
||||
}
|
||||
return out
|
||||
@@ -217,7 +234,7 @@ func (e *componentEncoder) encodeGroups() []*proto.GroupCompact {
|
||||
// list and a map from policy pointer to the indexes of its emitted rules in
|
||||
// that list — used by encodeResourcePoliciesMap to translate
|
||||
// ResourcePoliciesMap[resourceID][]*Policy into wire-side indexes.
|
||||
func (e *componentEncoder) encodePolicies(policies []*types.Policy) []*proto.PolicyCompact {
|
||||
func (e *componentEncoder) encodePolicies(policies []*nmdata.Policy) []*proto.PolicyCompact {
|
||||
if len(policies) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -239,7 +256,7 @@ func (e *componentEncoder) encodePolicies(policies []*types.Policy) []*proto.Pol
|
||||
}
|
||||
|
||||
// encodePolicyRule maps a single PolicyRule under pol to a PolicyCompact entry.
|
||||
func (e *componentEncoder) encodePolicyRule(pol *types.Policy, r *types.PolicyRule) *proto.PolicyCompact {
|
||||
func (e *componentEncoder) encodePolicyRule(pol *nmdata.Policy, r *nmdata.PolicyRule) *proto.PolicyCompact {
|
||||
return &proto.PolicyCompact{
|
||||
Id: pol.PublicID,
|
||||
Action: networkmap.GetProtoAction(string(r.Action)),
|
||||
@@ -278,14 +295,14 @@ func (e *componentEncoder) groupPublicXids(src []string) []string {
|
||||
// only live in ResourcePoliciesMap; without this union step they'd be lost
|
||||
// from the wire and the client's resource-policy lookup would come back
|
||||
// empty.
|
||||
func unionPolicies(policies []*types.Policy, resourcePolicies map[string][]*types.Policy) []*types.Policy {
|
||||
func unionPolicies(policies []*nmdata.Policy, resourcePolicies map[string][]*nmdata.Policy) []*nmdata.Policy {
|
||||
// Fast path: non-router peers have no resource-only policies, so the
|
||||
// "union" is identical to `policies`. Skip the dedup map allocation.
|
||||
if len(resourcePolicies) == 0 {
|
||||
return policies
|
||||
}
|
||||
seen := make(map[string]struct{}, len(policies))
|
||||
out := make([]*types.Policy, 0, len(policies))
|
||||
out := make([]*nmdata.Policy, 0, len(policies))
|
||||
for _, p := range policies {
|
||||
if p == nil {
|
||||
continue
|
||||
@@ -343,18 +360,31 @@ func (e *componentEncoder) groupPublicXid(groupID string) (string, bool) {
|
||||
// peers array. For other resource types only the type string is shipped
|
||||
// today (Calculate's resource-typed rule path consults SourceResource only
|
||||
// for "peer" — other types fall through to group-based lookup).
|
||||
func (e *componentEncoder) resourceToProto(r types.Resource) *proto.ResourceCompact {
|
||||
if r.ID == "" && r.Type == "" {
|
||||
func (e *componentEncoder) resourceToProto(r nmdata.Resource) *proto.ResourceCompact {
|
||||
t, ok := proto.ResourceCompactType_value[string(r.Type)]
|
||||
if !ok || t == 0 || r.ID == "" {
|
||||
return nil
|
||||
}
|
||||
out := &proto.ResourceCompact{Type: string(r.Type)}
|
||||
if r.Type == types.ResourceTypePeer && r.ID != "" {
|
||||
if idx, ok := e.peerOrder[r.ID]; ok {
|
||||
out.PeerIndexSet = true
|
||||
out.PeerIndex = idx
|
||||
if t == int32(proto.ResourceCompactType_peer) {
|
||||
idx, ok := e.peerOrder[r.ID]
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
return &proto.ResourceCompact{
|
||||
Type: proto.ResourceCompactType_peer,
|
||||
ResourceId: &proto.ResourceCompact_PeerIndex{PeerIndex: idx},
|
||||
}
|
||||
}
|
||||
return out
|
||||
|
||||
publicID, ok := e.networkIdToPublicId[r.ID]
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
return &proto.ResourceCompact{
|
||||
Type: proto.ResourceCompactType(t),
|
||||
ResourceId: &proto.ResourceCompact_Id{Id: publicID},
|
||||
}
|
||||
}
|
||||
|
||||
// postureCheckSeqs translates a slice of posture-check xids to their
|
||||
@@ -387,7 +417,7 @@ func (e *componentEncoder) networkPublicId(xid string) (string, bool) {
|
||||
return id, true
|
||||
}
|
||||
|
||||
func (e *componentEncoder) encodeDNSSettings(s *types.DNSSettings) *proto.DNSSettingsCompact {
|
||||
func (e *componentEncoder) encodeDNSSettings(s *nmdata.DNSSettings) *proto.DNSSettingsCompact {
|
||||
if s == nil || len(s.DisabledManagementGroups) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -402,7 +432,7 @@ func (e *componentEncoder) encodeDNSSettings(s *types.DNSSettings) *proto.DNSSet
|
||||
return out
|
||||
}
|
||||
|
||||
func (e *componentEncoder) encodeRoutes(routes []*nbroute.Route) []*proto.RouteRaw {
|
||||
func (e *componentEncoder) encodeRoutes(routes []*nmdata.Route) []*proto.RouteRaw {
|
||||
if len(routes) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -440,7 +470,7 @@ func (e *componentEncoder) encodeRoutes(routes []*nbroute.Route) []*proto.RouteR
|
||||
return out
|
||||
}
|
||||
|
||||
func (e *componentEncoder) encodeNameServerGroups(nsgs []*nbdns.NameServerGroup) []*proto.NameServerGroupRaw {
|
||||
func (e *componentEncoder) encodeNameServerGroups(nsgs []*nmdata.NameServerGroup) []*proto.NameServerGroupRaw {
|
||||
if len(nsgs) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -463,7 +493,7 @@ func (e *componentEncoder) encodeNameServerGroups(nsgs []*nbdns.NameServerGroup)
|
||||
return out
|
||||
}
|
||||
|
||||
func encodeNameServers(servers []nbdns.NameServer) []*proto.NameServer {
|
||||
func encodeNameServers(servers []nmdata.NameServer) []*proto.NameServer {
|
||||
if len(servers) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -478,7 +508,7 @@ func encodeNameServers(servers []nbdns.NameServer) []*proto.NameServer {
|
||||
return out
|
||||
}
|
||||
|
||||
func encodeSimpleRecords(records []nbdns.SimpleRecord) []*proto.SimpleRecord {
|
||||
func encodeSimpleRecords(records []nmdata.SimpleRecord) []*proto.SimpleRecord {
|
||||
if len(records) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -495,7 +525,7 @@ func encodeSimpleRecords(records []nbdns.SimpleRecord) []*proto.SimpleRecord {
|
||||
return out
|
||||
}
|
||||
|
||||
func encodeCustomZones(zones []nbdns.CustomZone) []*proto.CustomZone {
|
||||
func encodeCustomZones(zones []nmdata.CustomZone) []*proto.CustomZone {
|
||||
if len(zones) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -511,7 +541,7 @@ func encodeCustomZones(zones []nbdns.CustomZone) []*proto.CustomZone {
|
||||
return out
|
||||
}
|
||||
|
||||
func (e *componentEncoder) encodeNetworkResources(resources []*types.ComponentResource) []*proto.NetworkResourceRaw {
|
||||
func (e *componentEncoder) encodeNetworkResources(resources []*nmdata.NetworkResource) []*proto.NetworkResourceRaw {
|
||||
if len(resources) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -540,7 +570,7 @@ func (e *componentEncoder) encodeNetworkResources(resources []*types.ComponentRe
|
||||
return out
|
||||
}
|
||||
|
||||
func (e *componentEncoder) encodeRoutersMap(routersMap map[string]map[string]*types.ComponentRouter) map[string]*proto.NetworkRouterList {
|
||||
func (e *componentEncoder) encodeRoutersMap(routersMap map[string]map[string]*nmdata.NetworkRouter) map[string]*proto.NetworkRouterList {
|
||||
if len(routersMap) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -576,7 +606,7 @@ func (e *componentEncoder) encodeRoutersMap(routersMap map[string]map[string]*ty
|
||||
return out
|
||||
}
|
||||
|
||||
func (e *componentEncoder) encodeResourcePoliciesMap(rpm map[string][]*types.Policy) map[string]*proto.PolicyIds {
|
||||
func (e *componentEncoder) encodeResourcePoliciesMap(rpm map[string][]*nmdata.Policy) map[string]*proto.PolicyIds {
|
||||
if len(rpm) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -663,7 +693,7 @@ func (e *componentEncoder) encodePostureFailedPeers(m map[string]map[string]stru
|
||||
// (which shouldn't happen in production but the encoder is exported)
|
||||
// degrades to login_expiration_enabled = false, which makes
|
||||
// LoginExpired() return false for every peer.
|
||||
func toAccountSettingsCompact(s *types.AccountSettingsInfo) *proto.AccountSettingsCompact {
|
||||
func toAccountSettingsCompact(s *nmdata.AccountSettingsInfo) *proto.AccountSettingsCompact {
|
||||
if s == nil {
|
||||
return &proto.AccountSettingsCompact{}
|
||||
}
|
||||
@@ -673,7 +703,7 @@ func toAccountSettingsCompact(s *types.AccountSettingsInfo) *proto.AccountSettin
|
||||
}
|
||||
}
|
||||
|
||||
func toAccountNetwork(n *types.Network) *proto.AccountNetwork {
|
||||
func toAccountNetwork(n *nmdata.Network) *proto.AccountNetwork {
|
||||
if n == nil {
|
||||
return nil
|
||||
}
|
||||
@@ -689,20 +719,20 @@ func toAccountNetwork(n *types.Network) *proto.AccountNetwork {
|
||||
return out
|
||||
}
|
||||
|
||||
func toPeerCompact(p *types.ComponentPeer) *proto.PeerCompact {
|
||||
func toPeerCompact(p *nmdata.Peer) *proto.PeerCompact {
|
||||
pc := &proto.PeerCompact{
|
||||
WgPubKey: decodeWgKey(p.Key),
|
||||
SshPubKey: []byte(p.SSHKey),
|
||||
DnsLabel: p.DNSLabel,
|
||||
AgentVersion: p.AgentVersion,
|
||||
AddedWithSsoLogin: p.AddedWithSSOLogin,
|
||||
AgentVersion: p.Meta.WtVersion,
|
||||
AddedWithSsoLogin: p.UserID != "",
|
||||
LoginExpirationEnabled: p.LoginExpirationEnabled,
|
||||
SshEnabled: p.SSHEnabled,
|
||||
SupportsIpv6: p.SupportsIPv6,
|
||||
SupportsSourcePrefixes: p.SupportsSourcePrefixes,
|
||||
ServerSshAllowed: p.ServerSSHAllowed,
|
||||
SupportsIpv6: p.SupportsIPv6(),
|
||||
SupportsSourcePrefixes: p.SupportsSourcePrefixes(),
|
||||
ServerSshAllowed: p.Meta.Flags.ServerSSHAllowed,
|
||||
}
|
||||
if !p.LastLogin.IsZero() {
|
||||
if p.LastLogin != nil {
|
||||
pc.LastLoginUnixNano = p.LastLogin.UnixNano()
|
||||
}
|
||||
switch {
|
||||
@@ -751,7 +781,7 @@ func portsToUint32(ports []string) []uint32 {
|
||||
return out
|
||||
}
|
||||
|
||||
func portRangesToProto(ranges []types.RulePortRange) []*proto.PortInfo_Range {
|
||||
func portRangesToProto(ranges []nmdata.RulePortRange) []*proto.PortInfo_Range {
|
||||
if len(ranges) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -16,7 +16,7 @@ import (
|
||||
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
nbroute "github.com/netbirdio/netbird/route"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
@@ -152,66 +152,66 @@ func envelopesEquivalent(a, b *proto.NetworkMapEnvelope) bool {
|
||||
}
|
||||
|
||||
func newTestComponents() *types.NetworkMapComponents {
|
||||
peerA := &types.ComponentPeer{
|
||||
ID: "peer-a",
|
||||
Key: testWgKeyA,
|
||||
IP: netip.AddrFrom4([4]byte{100, 64, 0, 1}),
|
||||
DNSLabel: "peera",
|
||||
SSHKey: "ssh-a",
|
||||
AgentVersion: "0.40.0",
|
||||
peerA := &nmdata.Peer{
|
||||
ID: "peer-a",
|
||||
Key: testWgKeyA,
|
||||
IP: netip.AddrFrom4([4]byte{100, 64, 0, 1}),
|
||||
DNSLabel: "peera",
|
||||
SSHKey: "ssh-a",
|
||||
Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
|
||||
}
|
||||
peerB := &types.ComponentPeer{
|
||||
ID: "peer-b",
|
||||
Key: testWgKeyB,
|
||||
IP: netip.AddrFrom4([4]byte{100, 64, 0, 2}),
|
||||
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2}),
|
||||
DNSLabel: "peerb",
|
||||
AgentVersion: "0.25.0",
|
||||
peerB := &nmdata.Peer{
|
||||
ID: "peer-b",
|
||||
Key: testWgKeyB,
|
||||
IP: netip.AddrFrom4([4]byte{100, 64, 0, 2}),
|
||||
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2}),
|
||||
DNSLabel: "peerb",
|
||||
Meta: nmdata.PeerSystemMeta{WtVersion: "0.25.0"},
|
||||
}
|
||||
peerC := &types.ComponentPeer{
|
||||
ID: "peer-c",
|
||||
Key: testWgKeyC,
|
||||
IP: netip.AddrFrom4([4]byte{100, 64, 0, 3}),
|
||||
DNSLabel: "peerc",
|
||||
AgentVersion: "0.40.0",
|
||||
peerC := &nmdata.Peer{
|
||||
ID: "peer-c",
|
||||
Key: testWgKeyC,
|
||||
IP: netip.AddrFrom4([4]byte{100, 64, 0, 3}),
|
||||
DNSLabel: "peerc",
|
||||
Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
|
||||
}
|
||||
|
||||
return &types.NetworkMapComponents{
|
||||
PeerID: "peer-a",
|
||||
Network: &types.Network{
|
||||
Network: &nmdata.Network{
|
||||
Identifier: "net-test",
|
||||
Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)},
|
||||
Serial: 7,
|
||||
},
|
||||
AccountSettings: &types.AccountSettingsInfo{
|
||||
AccountSettings: &nmdata.AccountSettingsInfo{
|
||||
PeerLoginExpirationEnabled: true,
|
||||
PeerLoginExpiration: 2 * time.Hour,
|
||||
},
|
||||
Peers: map[string]*types.ComponentPeer{
|
||||
Peers: map[string]*nmdata.Peer{
|
||||
"peer-a": peerA,
|
||||
"peer-b": peerB,
|
||||
"peer-c": peerC,
|
||||
},
|
||||
Groups: map[string]*types.ComponentGroup{
|
||||
"group-src": {ID: "group-src", PublicID: "1", Name: "Src", Peers: []string{"peer-a"}},
|
||||
"group-dst": {ID: "group-dst", PublicID: "2", Name: "Dst", Peers: []string{"peer-b", "peer-c"}},
|
||||
Groups: map[string]*nmdata.Group{
|
||||
"group-src": {PublicID: "1", Name: "Src", Peers: []string{"peer-a"}},
|
||||
"group-dst": {PublicID: "2", Name: "Dst", Peers: []string{"peer-b", "peer-c"}},
|
||||
},
|
||||
Policies: []*types.Policy{
|
||||
Policies: []*nmdata.Policy{
|
||||
{
|
||||
ID: "pol-1",
|
||||
PublicID: "10",
|
||||
Enabled: true,
|
||||
Rules: []*types.PolicyRule{{
|
||||
ID: "rule-1", Enabled: true, Action: types.PolicyTrafficActionAccept,
|
||||
Protocol: types.PolicyRuleProtocolTCP, Bidirectional: true,
|
||||
Rules: []*nmdata.PolicyRule{{
|
||||
ID: "rule-1", Enabled: true, Action: string(types.PolicyTrafficActionAccept),
|
||||
Protocol: string(types.PolicyRuleProtocolTCP), Bidirectional: true,
|
||||
Ports: []string{"22", "80"},
|
||||
PortRanges: []types.RulePortRange{{Start: 8000, End: 8100}},
|
||||
PortRanges: []nmdata.RulePortRange{{Start: 8000, End: 8100}},
|
||||
Sources: []string{"group-src"},
|
||||
Destinations: []string{"group-dst"},
|
||||
}},
|
||||
},
|
||||
},
|
||||
RouterPeers: map[string]*types.ComponentPeer{"peer-c": peerC},
|
||||
RouterPeers: map[string]*nmdata.Peer{"peer-c": peerC},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -304,6 +304,31 @@ func TestEncodeNetworkMapEnvelope_GroupsByAccountPublicId(t *testing.T) {
|
||||
assert.Len(t, groupByID["2"].PeerIndexes, 2)
|
||||
}
|
||||
|
||||
func TestEncodePolicy(t *testing.T) {
|
||||
encoder := componentEncoder{peerOrder: map[string]uint32{"peerId": uint32(1234)}, networkIdToPublicId: map[string]string{"domain": "publicDomain", "host": "publicHost", "subnet": "publicSubnet"}}
|
||||
assert.Equal(t,
|
||||
encoder.resourceToProto(nmdata.Resource{Type: "peer", ID: "peerId"}),
|
||||
&proto.ResourceCompact{Type: proto.ResourceCompactType_peer, ResourceId: &proto.ResourceCompact_PeerIndex{PeerIndex: uint32(1234)}})
|
||||
// verify invalid peer id results in nil
|
||||
assert.Nil(t,
|
||||
encoder.resourceToProto(nmdata.Resource{Type: "peer", ID: "boom"}))
|
||||
assert.Equal(t,
|
||||
encoder.resourceToProto(nmdata.Resource{Type: "domain", ID: "domain"}),
|
||||
&proto.ResourceCompact{Type: proto.ResourceCompactType_domain, ResourceId: &proto.ResourceCompact_Id{Id: "publicDomain"}})
|
||||
assert.Equal(t,
|
||||
encoder.resourceToProto(nmdata.Resource{Type: "host", ID: "host"}),
|
||||
&proto.ResourceCompact{Type: proto.ResourceCompactType_host, ResourceId: &proto.ResourceCompact_Id{Id: "publicHost"}})
|
||||
assert.Equal(t,
|
||||
encoder.resourceToProto(nmdata.Resource{Type: "subnet", ID: "subnet"}),
|
||||
&proto.ResourceCompact{Type: proto.ResourceCompactType_subnet, ResourceId: &proto.ResourceCompact_Id{Id: "publicSubnet"}})
|
||||
// verify invalid resource type results in nil
|
||||
assert.Nil(t,
|
||||
encoder.resourceToProto(nmdata.Resource{Type: "boom", ID: "boom"}))
|
||||
// verify invalid networkresource id results in nil
|
||||
assert.Nil(t,
|
||||
encoder.resourceToProto(nmdata.Resource{Type: "host", ID: "boom"}))
|
||||
}
|
||||
|
||||
func TestEncodeNetworkMapEnvelope_PolicyExpansion(t *testing.T) {
|
||||
c := newTestComponents()
|
||||
|
||||
@@ -377,12 +402,12 @@ func TestEncodeNetworkMapEnvelope_MalformedWgKey(t *testing.T) {
|
||||
|
||||
func TestEncodeNetworkMapEnvelope_IPv6OnlyPeer(t *testing.T) {
|
||||
c := newTestComponents()
|
||||
v6Only := &types.ComponentPeer{
|
||||
ID: "peer-v6",
|
||||
Key: testWgKeyA,
|
||||
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 9}),
|
||||
DNSLabel: "peerv6",
|
||||
AgentVersion: "0.40.0",
|
||||
v6Only := &nmdata.Peer{
|
||||
ID: "peer-v6",
|
||||
Key: testWgKeyA,
|
||||
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 9}),
|
||||
DNSLabel: "peerv6",
|
||||
Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
|
||||
}
|
||||
c.Peers["peer-v6"] = v6Only
|
||||
|
||||
@@ -401,11 +426,11 @@ func TestEncodeNetworkMapEnvelope_IPv6OnlyPeer(t *testing.T) {
|
||||
|
||||
func TestEncodeNetworkMapEnvelope_PeerWithoutIP(t *testing.T) {
|
||||
c := newTestComponents()
|
||||
c.Peers["peer-noip"] = &types.ComponentPeer{
|
||||
ID: "peer-noip",
|
||||
Key: testWgKeyA,
|
||||
DNSLabel: "peernoip",
|
||||
AgentVersion: "0.40.0",
|
||||
c.Peers["peer-noip"] = &nmdata.Peer{
|
||||
ID: "peer-noip",
|
||||
Key: testWgKeyA,
|
||||
DNSLabel: "peernoip",
|
||||
Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
|
||||
}
|
||||
|
||||
full := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c}).GetFull()
|
||||
@@ -423,7 +448,7 @@ func TestEncodeNetworkMapEnvelope_PeerWithoutIP(t *testing.T) {
|
||||
|
||||
func TestEncodeNetworkMapEnvelope_EmptyInput(t *testing.T) {
|
||||
c := &types.NetworkMapComponents{
|
||||
Network: &types.Network{Identifier: "x", Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}},
|
||||
Network: &nmdata.Network{Identifier: "x", Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}},
|
||||
}
|
||||
|
||||
env := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c})
|
||||
@@ -440,9 +465,9 @@ func TestEncodeNetworkMapEnvelope_EmptyInput(t *testing.T) {
|
||||
func TestEncodeNetworkMapEnvelope_PeerLoginExpirationFields(t *testing.T) {
|
||||
c := newTestComponents()
|
||||
now := time.Date(2024, 1, 2, 3, 4, 5, 0, time.UTC)
|
||||
c.Peers["peer-a"].AddedWithSSOLogin = true
|
||||
c.Peers["peer-a"].UserID = "user-1"
|
||||
c.Peers["peer-a"].LoginExpirationEnabled = true
|
||||
c.Peers["peer-a"].LastLogin = now
|
||||
c.Peers["peer-a"].LastLogin = &now
|
||||
|
||||
full := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c}).GetFull()
|
||||
|
||||
@@ -472,7 +497,7 @@ func TestEncodeNetworkMapEnvelope_PeerLoginExpirationFields(t *testing.T) {
|
||||
|
||||
func TestEncodeNetworkMapEnvelope_RoutesRoundTrip(t *testing.T) {
|
||||
c := newTestComponents()
|
||||
c.Routes = []*nbroute.Route{
|
||||
c.Routes = []*nmdata.Route{
|
||||
{
|
||||
ID: "route-peer",
|
||||
PublicID: "100",
|
||||
@@ -519,7 +544,7 @@ func TestEncodeNetworkMapEnvelope_RoutesRoundTrip(t *testing.T) {
|
||||
|
||||
func TestEncodeNetworkMapEnvelope_RouteWithMissingPeerLeavesIndexUnset(t *testing.T) {
|
||||
c := newTestComponents()
|
||||
c.Routes = []*nbroute.Route{{
|
||||
c.Routes = []*nmdata.Route{{
|
||||
ID: "route-x",
|
||||
PublicID: "100",
|
||||
Peer: "peer-not-in-components",
|
||||
@@ -539,21 +564,21 @@ func TestEncodeNetworkMapEnvelope_ResourceOnlyPolicyShippedAndIndexed(t *testing
|
||||
// Policy that exists ONLY in ResourcePoliciesMap, not in c.Policies. This
|
||||
// is the I1 case — without unionPolicies the encoder would silently
|
||||
// drop it from the wire.
|
||||
resourceOnlyPolicy := &types.Policy{
|
||||
resourceOnlyPolicy := &nmdata.Policy{
|
||||
ID: "pol-resource", PublicID: "99", Enabled: true,
|
||||
Rules: []*types.PolicyRule{{
|
||||
ID: "rule-r", Enabled: true, Action: types.PolicyTrafficActionAccept,
|
||||
Protocol: types.PolicyRuleProtocolTCP,
|
||||
Rules: []*nmdata.PolicyRule{{
|
||||
ID: "rule-r", Enabled: true, Action: string(types.PolicyTrafficActionAccept),
|
||||
Protocol: string(types.PolicyRuleProtocolTCP),
|
||||
Sources: []string{"group-src"},
|
||||
Destinations: []string{"group-dst"},
|
||||
}},
|
||||
}
|
||||
c.ResourcePoliciesMap = map[string][]*types.Policy{
|
||||
c.ResourcePoliciesMap = map[string][]*nmdata.Policy{
|
||||
"resource-x": {c.Policies[0], resourceOnlyPolicy}, // shared + resource-only
|
||||
}
|
||||
// Resource must appear in components.NetworkResources with a seq id —
|
||||
// encoder uses that to translate the xid map key to uint32.
|
||||
c.NetworkResources = []*types.ComponentResource{
|
||||
c.NetworkResources = []*nmdata.NetworkResource{
|
||||
{ID: "resource-x", PublicID: "77", Name: "res-x", Enabled: true},
|
||||
}
|
||||
|
||||
@@ -562,27 +587,16 @@ func TestEncodeNetworkMapEnvelope_ResourceOnlyPolicyShippedAndIndexed(t *testing
|
||||
require.Len(t, full.Policies, 2, "encoded policies must include both peer-traffic and resource-only")
|
||||
|
||||
policyByID := map[string]*proto.PolicyCompact{}
|
||||
policyIds := make([]string, 0)
|
||||
for _, p := range full.Policies {
|
||||
policyByID[p.Id] = p
|
||||
policyIds = append(policyIds, p.Id)
|
||||
}
|
||||
require.Contains(t, policyByID, "10", "original peer-traffic policy id 10")
|
||||
require.Contains(t, policyByID, "99", "resource-only policy id 99")
|
||||
|
||||
require.Contains(t, full.ResourcePoliciesMap, "77")
|
||||
ids := full.ResourcePoliciesMap["77"].Ids
|
||||
require.Len(t, ids, 2)
|
||||
assert.ElementsMatch(t, policyIds, ids,
|
||||
"resource policies map must reference both wire policy indexes")
|
||||
}
|
||||
|
||||
func TestEncodeNetworkMapEnvelope_NameServerGroups(t *testing.T) {
|
||||
c := newTestComponents()
|
||||
c.NameServerGroups = []*nbdns.NameServerGroup{{
|
||||
c.NameServerGroups = []*nmdata.NameServerGroup{{
|
||||
ID: "nsg-1", PublicID: "50", Name: "Main", Description: "primary",
|
||||
NameServers: []nbdns.NameServer{{
|
||||
IP: netip.MustParseAddr("8.8.8.8"), NSType: nbdns.UDPNameServerType, Port: 53,
|
||||
NameServers: []nmdata.NameServer{{
|
||||
IP: netip.MustParseAddr("8.8.8.8"), NSType: int(nbdns.UDPNameServerType), Port: 53,
|
||||
}},
|
||||
Groups: []string{"group-src", "group-not-persisted"},
|
||||
Primary: true, Enabled: true,
|
||||
@@ -621,11 +635,11 @@ func TestEncodeNetworkMapEnvelope_PostureFailedPeers(t *testing.T) {
|
||||
func TestEncodeNetworkMapEnvelope_RoutersMap(t *testing.T) {
|
||||
c := newTestComponents()
|
||||
c.NetworkXIDToPublicID = map[string]string{"net-1": "5"}
|
||||
c.RoutersMap = map[string]map[string]*types.ComponentRouter{
|
||||
c.RoutersMap = map[string]map[string]*nmdata.NetworkRouter{
|
||||
"net-1": {
|
||||
"peer-c": {
|
||||
PublicID: "200",
|
||||
Peer: "peer-c", Masquerade: true, Metric: 10, Enabled: true,
|
||||
PublicID: "200",
|
||||
Masquerade: true, Metric: 10, Enabled: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
@@ -651,14 +665,14 @@ func TestEncodeNetworkMapEnvelope_RouterPeerNotInComponentsPeers(t *testing.T) {
|
||||
// peer_index reference must still resolve.
|
||||
c := newTestComponents()
|
||||
delete(c.Peers, "peer-c")
|
||||
routerPeer := &types.ComponentPeer{
|
||||
routerPeer := &nmdata.Peer{
|
||||
ID: "peer-c", Key: testWgKeyC, IP: netip.AddrFrom4([4]byte{100, 64, 0, 3}),
|
||||
DNSLabel: "peerc", AgentVersion: "0.40.0",
|
||||
DNSLabel: "peerc", Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
|
||||
}
|
||||
c.RouterPeers = map[string]*types.ComponentPeer{"peer-c": routerPeer}
|
||||
c.RouterPeers = map[string]*nmdata.Peer{"peer-c": routerPeer}
|
||||
c.NetworkXIDToPublicID = map[string]string{"net-1": "5"}
|
||||
c.RoutersMap = map[string]map[string]*types.ComponentRouter{
|
||||
"net-1": {"peer-c": {PublicID: "1", Peer: "peer-c", Enabled: true}},
|
||||
c.RoutersMap = map[string]map[string]*nmdata.NetworkRouter{
|
||||
"net-1": {"peer-c": {PublicID: "1", Enabled: true}},
|
||||
}
|
||||
|
||||
full := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c}).GetFull()
|
||||
@@ -691,9 +705,9 @@ func TestToProxyPatch_EmptyInputReturnsNil(t *testing.T) {
|
||||
|
||||
func TestToProxyPatch_PopulatesAllFields(t *testing.T) {
|
||||
nm := &types.NetworkMap{
|
||||
Peers: []*types.ComponentPeer{{
|
||||
Peers: []*nmdata.Peer{{
|
||||
ID: "ext-peer", Key: testWgKeyA, IP: netip.AddrFrom4([4]byte{100, 64, 0, 9}),
|
||||
DNSLabel: "extpeer", AgentVersion: "0.40.0",
|
||||
DNSLabel: "extpeer", Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
|
||||
}},
|
||||
FirewallRules: []*types.FirewallRule{{
|
||||
PeerIP: "100.64.0.9", Action: "accept", Direction: 0, Protocol: "tcp",
|
||||
@@ -762,7 +776,7 @@ func TestEncodeNetworkMapEnvelope_NilComponentsGracefulDegrade(t *testing.T) {
|
||||
|
||||
func TestEncodeNetworkMapEnvelope_AccountSettingsAlwaysEmitted(t *testing.T) {
|
||||
c := &types.NetworkMapComponents{
|
||||
Network: &types.Network{Identifier: "x", Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}},
|
||||
Network: &nmdata.Network{Identifier: "x", Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}},
|
||||
// AccountSettings deliberately nil
|
||||
}
|
||||
|
||||
@@ -776,6 +790,6 @@ func TestEncodeNetworkMapEnvelope_AccountSettingsAlwaysEmitted(t *testing.T) {
|
||||
func emptyNetworkMapComponents() *types.NetworkMapComponents {
|
||||
return types.EmptyNetworkMapComponents(
|
||||
&types.NetworkMapComponents{
|
||||
PeerID: "peer-id", Peers: map[string]*types.ComponentPeer{"peer-id": {}}},
|
||||
PeerID: "peer-id", Peers: map[string]*nmdata.Peer{"peer-id": {}}},
|
||||
)
|
||||
}
|
||||
|
||||
@@ -7,11 +7,11 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/client/ssh/auth"
|
||||
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/posture"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
sharedgrpc "github.com/netbirdio/netbird/shared/management/grpc"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap"
|
||||
nmdata "github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
@@ -31,14 +31,14 @@ func ToComponentSyncResponse(
|
||||
config *nbconfig.Config,
|
||||
httpConfig *nbconfig.HttpServerConfig,
|
||||
deviceFlowConfig *nbconfig.DeviceAuthorizationFlow,
|
||||
peer *nbpeer.Peer,
|
||||
peer *nmdata.Peer,
|
||||
turnCredentials *Token,
|
||||
relayCredentials *Token,
|
||||
components *types.NetworkMapComponents,
|
||||
proxyPatch *types.NetworkMap,
|
||||
dnsName string,
|
||||
checks []*posture.Checks,
|
||||
settings *types.Settings,
|
||||
settings *nmdata.AccountSettingsInfo,
|
||||
extraSettings *types.ExtraSettings,
|
||||
peerGroups []string,
|
||||
dnsFwdPort int64,
|
||||
@@ -50,7 +50,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()
|
||||
@@ -145,7 +145,7 @@ func toProxyPatch(nm *types.NetworkMap, dnsName string, includeIPv6, useSourcePr
|
||||
//
|
||||
// The full SSH AuthorizedUsers map is still produced by the client when it
|
||||
// runs Calculate() over the envelope.
|
||||
func computeSSHEnabledForPeer(c *types.NetworkMapComponents, peer *nbpeer.Peer) bool {
|
||||
func computeSSHEnabledForPeer(c *types.NetworkMapComponents, peer *nmdata.Peer) bool {
|
||||
if c == nil || peer == nil {
|
||||
return false
|
||||
}
|
||||
@@ -170,25 +170,25 @@ func computeSSHEnabledForPeer(c *types.NetworkMapComponents, peer *nbpeer.Peer)
|
||||
// ruleEnablesSSHForPeer returns true when rule is active, targets peer, and
|
||||
// either explicitly authorises SSH or covers the legacy TCP/22 path while the
|
||||
// peer itself has SSH enabled locally.
|
||||
func ruleEnablesSSHForPeer(c *types.NetworkMapComponents, rule *types.PolicyRule, peer *nbpeer.Peer) bool {
|
||||
func ruleEnablesSSHForPeer(c *types.NetworkMapComponents, rule *nmdata.PolicyRule, peer *nmdata.Peer) bool {
|
||||
if rule == nil || !rule.Enabled {
|
||||
return false
|
||||
}
|
||||
if !peerInDestinations(c, rule, peer.ID) {
|
||||
return false
|
||||
}
|
||||
if rule.Protocol == types.PolicyRuleProtocolNetbirdSSH {
|
||||
if rule.Protocol == string(types.PolicyRuleProtocolNetbirdSSH) {
|
||||
return true
|
||||
}
|
||||
return peer.SSHEnabled && types.PolicyRuleImpliesLegacySSH(rule)
|
||||
return peer.SSHEnabled && nmdata.PolicyRuleImpliesLegacySSH(rule)
|
||||
}
|
||||
|
||||
// peerInDestinations reports whether peerID is in any of rule.Destinations'
|
||||
// groups (or matches DestinationResource if it's a peer-typed resource —
|
||||
// for non-peer types Calculate falls through to group lookup, so we mirror
|
||||
// that exactly to avoid silent divergence).
|
||||
func peerInDestinations(c *types.NetworkMapComponents, rule *types.PolicyRule, peerID string) bool {
|
||||
if rule.DestinationResource.Type == types.ResourceTypePeer && rule.DestinationResource.ID != "" {
|
||||
func peerInDestinations(c *types.NetworkMapComponents, rule *nmdata.PolicyRule, peerID string) bool {
|
||||
if rule.DestinationResource.Type == string(types.ResourceTypePeer) && rule.DestinationResource.ID != "" {
|
||||
return rule.DestinationResource.ID == peerID
|
||||
}
|
||||
for _, groupID := range rule.Destinations {
|
||||
|
||||
@@ -5,8 +5,8 @@ import (
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
// TestComputeSSHEnabledForPeer covers both Calculate-mirroring branches:
|
||||
@@ -17,16 +17,15 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
|
||||
const targetPeerID = "target"
|
||||
const targetGroupID = "g_dst"
|
||||
|
||||
mkComponents := func(rule *types.PolicyRule, sshEnabled bool) (*types.NetworkMapComponents, *nbpeer.Peer) {
|
||||
peer := &nbpeer.Peer{ID: targetPeerID, SSHEnabled: sshEnabled}
|
||||
group := &types.ComponentGroup{ID: targetGroupID, Name: "dst", Peers: []string{targetPeerID}}
|
||||
mkComponents := func(rule *nmdata.PolicyRule, sshEnabled bool) (*types.NetworkMapComponents, *nmdata.Peer) {
|
||||
peer := &nmdata.Peer{ID: targetPeerID, SSHEnabled: sshEnabled}
|
||||
return &types.NetworkMapComponents{
|
||||
Peers: map[string]*types.ComponentPeer{targetPeerID: peer.ToComponent()},
|
||||
Groups: map[string]*types.ComponentGroup{targetGroupID: group},
|
||||
Policies: []*types.Policy{{
|
||||
Peers: map[string]*nmdata.Peer{targetPeerID: peer},
|
||||
Groups: map[string]*nmdata.Group{targetGroupID: {Name: "dst", Peers: []string{targetPeerID}}},
|
||||
Policies: []*nmdata.Policy{{
|
||||
ID: "p",
|
||||
Enabled: true,
|
||||
Rules: []*types.PolicyRule{rule},
|
||||
Rules: []*nmdata.PolicyRule{rule},
|
||||
}},
|
||||
}, peer
|
||||
}
|
||||
@@ -34,14 +33,14 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
peerSSH bool
|
||||
rule types.PolicyRule
|
||||
rule nmdata.PolicyRule
|
||||
wantEnabled bool
|
||||
}{
|
||||
{
|
||||
name: "explicit-netbird-ssh-activates-regardless-of-peer-ssh",
|
||||
peerSSH: false,
|
||||
rule: types.PolicyRule{
|
||||
Enabled: true, Protocol: types.PolicyRuleProtocolNetbirdSSH,
|
||||
rule: nmdata.PolicyRule{
|
||||
Enabled: true, Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
|
||||
Destinations: []string{targetGroupID},
|
||||
},
|
||||
wantEnabled: true,
|
||||
@@ -49,8 +48,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
|
||||
{
|
||||
name: "implicit-tcp-22-with-peer-ssh",
|
||||
peerSSH: true,
|
||||
rule: types.PolicyRule{
|
||||
Enabled: true, Protocol: types.PolicyRuleProtocolTCP, Ports: []string{"22"},
|
||||
rule: nmdata.PolicyRule{
|
||||
Enabled: true, Protocol: string(types.PolicyRuleProtocolTCP), Ports: []string{"22"},
|
||||
Destinations: []string{targetGroupID},
|
||||
},
|
||||
wantEnabled: true,
|
||||
@@ -58,8 +57,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
|
||||
{
|
||||
name: "implicit-tcp-22-without-peer-ssh-disabled",
|
||||
peerSSH: false,
|
||||
rule: types.PolicyRule{
|
||||
Enabled: true, Protocol: types.PolicyRuleProtocolTCP, Ports: []string{"22"},
|
||||
rule: nmdata.PolicyRule{
|
||||
Enabled: true, Protocol: string(types.PolicyRuleProtocolTCP), Ports: []string{"22"},
|
||||
Destinations: []string{targetGroupID},
|
||||
},
|
||||
wantEnabled: false,
|
||||
@@ -67,8 +66,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
|
||||
{
|
||||
name: "implicit-tcp-22022-with-peer-ssh",
|
||||
peerSSH: true,
|
||||
rule: types.PolicyRule{
|
||||
Enabled: true, Protocol: types.PolicyRuleProtocolTCP, Ports: []string{"22022"},
|
||||
rule: nmdata.PolicyRule{
|
||||
Enabled: true, Protocol: string(types.PolicyRuleProtocolTCP), Ports: []string{"22022"},
|
||||
Destinations: []string{targetGroupID},
|
||||
},
|
||||
wantEnabled: true,
|
||||
@@ -76,8 +75,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
|
||||
{
|
||||
name: "implicit-all-protocol-with-peer-ssh",
|
||||
peerSSH: true,
|
||||
rule: types.PolicyRule{
|
||||
Enabled: true, Protocol: types.PolicyRuleProtocolALL,
|
||||
rule: nmdata.PolicyRule{
|
||||
Enabled: true, Protocol: string(types.PolicyRuleProtocolALL),
|
||||
Destinations: []string{targetGroupID},
|
||||
},
|
||||
wantEnabled: true,
|
||||
@@ -85,10 +84,10 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
|
||||
{
|
||||
name: "implicit-port-range-covers-22",
|
||||
peerSSH: true,
|
||||
rule: types.PolicyRule{
|
||||
rule: nmdata.PolicyRule{
|
||||
Enabled: true,
|
||||
Protocol: types.PolicyRuleProtocolTCP,
|
||||
PortRanges: []types.RulePortRange{{Start: 20, End: 30}},
|
||||
Protocol: string(types.PolicyRuleProtocolTCP),
|
||||
PortRanges: []nmdata.RulePortRange{{Start: 20, End: 30}},
|
||||
Destinations: []string{targetGroupID},
|
||||
},
|
||||
wantEnabled: true,
|
||||
@@ -96,8 +95,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
|
||||
{
|
||||
name: "tcp-80-no-ssh",
|
||||
peerSSH: true,
|
||||
rule: types.PolicyRule{
|
||||
Enabled: true, Protocol: types.PolicyRuleProtocolTCP, Ports: []string{"80"},
|
||||
rule: nmdata.PolicyRule{
|
||||
Enabled: true, Protocol: string(types.PolicyRuleProtocolTCP), Ports: []string{"80"},
|
||||
Destinations: []string{targetGroupID},
|
||||
},
|
||||
wantEnabled: false,
|
||||
@@ -105,8 +104,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
|
||||
{
|
||||
name: "disabled-rule-skipped",
|
||||
peerSSH: true,
|
||||
rule: types.PolicyRule{
|
||||
Enabled: false, Protocol: types.PolicyRuleProtocolNetbirdSSH,
|
||||
rule: nmdata.PolicyRule{
|
||||
Enabled: false, Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
|
||||
Destinations: []string{targetGroupID},
|
||||
},
|
||||
wantEnabled: false,
|
||||
@@ -114,8 +113,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
|
||||
{
|
||||
name: "peer-not-in-destinations",
|
||||
peerSSH: true,
|
||||
rule: types.PolicyRule{
|
||||
Enabled: true, Protocol: types.PolicyRuleProtocolNetbirdSSH,
|
||||
rule: nmdata.PolicyRule{
|
||||
Enabled: true, Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
|
||||
Destinations: []string{"g_other"}, // target not in this group
|
||||
},
|
||||
wantEnabled: false,
|
||||
@@ -123,21 +122,21 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
|
||||
{
|
||||
name: "peer-typed-destination-resource-matches",
|
||||
peerSSH: false,
|
||||
rule: types.PolicyRule{
|
||||
rule: nmdata.PolicyRule{
|
||||
Enabled: true,
|
||||
Protocol: types.PolicyRuleProtocolNetbirdSSH,
|
||||
DestinationResource: types.Resource{ID: targetPeerID, Type: types.ResourceTypePeer},
|
||||
Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
|
||||
DestinationResource: nmdata.Resource{ID: targetPeerID, Type: string(types.ResourceTypePeer)},
|
||||
},
|
||||
wantEnabled: true,
|
||||
},
|
||||
{
|
||||
name: "non-peer-destination-resource-falls-through-to-groups",
|
||||
peerSSH: false,
|
||||
rule: types.PolicyRule{
|
||||
rule: nmdata.PolicyRule{
|
||||
Enabled: true,
|
||||
Protocol: types.PolicyRuleProtocolNetbirdSSH,
|
||||
DestinationResource: types.Resource{ID: targetPeerID, Type: "host"}, // wrong type
|
||||
Destinations: []string{targetGroupID}, // saved by group fallback
|
||||
Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
|
||||
DestinationResource: nmdata.Resource{ID: targetPeerID, Type: "host"}, // wrong type
|
||||
Destinations: []string{targetGroupID}, // saved by group fallback
|
||||
},
|
||||
wantEnabled: true,
|
||||
},
|
||||
@@ -156,16 +155,16 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
|
||||
// belt-and-suspenders presence guard mirroring Calculate's
|
||||
// getAllPeersFromGroups invariant.
|
||||
func TestComputeSSHEnabledForPeer_TargetMissingFromComponents(t *testing.T) {
|
||||
peer := &nbpeer.Peer{ID: "missing", SSHEnabled: true}
|
||||
peer := &nmdata.Peer{ID: "missing", SSHEnabled: true}
|
||||
c := &types.NetworkMapComponents{
|
||||
Peers: map[string]*types.ComponentPeer{}, // target peer NOT present
|
||||
Groups: map[string]*types.ComponentGroup{
|
||||
"g": {ID: "g", Peers: []string{"missing"}},
|
||||
Peers: map[string]*nmdata.Peer{}, // target peer NOT present
|
||||
Groups: map[string]*nmdata.Group{
|
||||
"g": {Peers: []string{"missing"}},
|
||||
},
|
||||
Policies: []*types.Policy{{
|
||||
Policies: []*nmdata.Policy{{
|
||||
ID: "p", Enabled: true,
|
||||
Rules: []*types.PolicyRule{{
|
||||
Enabled: true, Protocol: types.PolicyRuleProtocolNetbirdSSH,
|
||||
Rules: []*nmdata.PolicyRule{{
|
||||
Enabled: true, Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
|
||||
Destinations: []string{"g"},
|
||||
}},
|
||||
}},
|
||||
@@ -179,6 +178,6 @@ func TestComputeSSHEnabledForPeer_TargetMissingFromComponents(t *testing.T) {
|
||||
// exported indirectly via ToComponentSyncResponse and may receive nil
|
||||
// components on graceful-degrade paths.
|
||||
func TestComputeSSHEnabledForPeer_NilInputs(t *testing.T) {
|
||||
assert.False(t, computeSSHEnabledForPeer(nil, &nbpeer.Peer{ID: "x"}))
|
||||
assert.False(t, computeSSHEnabledForPeer(nil, &nmdata.Peer{ID: "x"}))
|
||||
assert.False(t, computeSSHEnabledForPeer(&types.NetworkMapComponents{}, nil))
|
||||
}
|
||||
|
||||
@@ -18,10 +18,10 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map/controller/cache"
|
||||
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/posture"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
"github.com/netbirdio/netbird/shared/netiputil"
|
||||
)
|
||||
@@ -47,7 +47,7 @@ func init() {
|
||||
// nil when no server config is set (the fan-out network-map path) because clients treat any
|
||||
// non-nil config as authoritative: a config without a relay section is interpreted as relay
|
||||
// disabled and wipes the clients' relay URLs.
|
||||
func toNetbirdConfig(config *nbconfig.Config, turnCredentials *Token, relayToken *Token, extraSettings *types.ExtraSettings, settings *types.Settings) *proto.NetbirdConfig {
|
||||
func toNetbirdConfig(config *nbconfig.Config, turnCredentials *Token, relayToken *Token, extraSettings *types.ExtraSettings, settings *nmdata.AccountSettingsInfo) *proto.NetbirdConfig {
|
||||
if config == nil {
|
||||
return nil
|
||||
}
|
||||
@@ -119,7 +119,7 @@ func toNetbirdConfig(config *nbconfig.Config, turnCredentials *Token, relayToken
|
||||
return nbConfig
|
||||
}
|
||||
|
||||
func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, settings *types.Settings, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, enableSSH bool) *proto.PeerConfig {
|
||||
func toPeerConfig(peer *nmdata.Peer, network *nmdata.Network, dnsName string, settings *nmdata.AccountSettingsInfo, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, enableSSH bool, forceRoutingPeerDNS bool) *proto.PeerConfig {
|
||||
netmask, _ := network.Net.Mask.Size()
|
||||
fqdn := peer.FQDN(dnsName)
|
||||
|
||||
@@ -135,7 +135,7 @@ func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, set
|
||||
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,
|
||||
@@ -154,7 +154,7 @@ func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, set
|
||||
return peerConfig
|
||||
}
|
||||
|
||||
func ToSyncResponse(ctx context.Context, config *nbconfig.Config, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, peer *nbpeer.Peer, turnCredentials *Token, relayCredentials *Token, networkMap *types.NetworkMap, dnsName string, checks []*posture.Checks, dnsCache *cache.DNSConfigCache, settings *types.Settings, extraSettings *types.ExtraSettings, peerGroups []string, dnsFwdPort int64) *proto.SyncResponse {
|
||||
func ToSyncResponse(ctx context.Context, config *nbconfig.Config, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, peer *nmdata.Peer, turnCredentials *Token, relayCredentials *Token, networkMap *types.NetworkMap, dnsName string, checks []*posture.Checks, dnsCache *cache.DNSConfigCache, settings *nmdata.AccountSettingsInfo, extraSettings *types.ExtraSettings, peerGroups []string, dnsFwdPort int64) *proto.SyncResponse {
|
||||
// IPv6 data in AllowedIPs and SourcePrefixes wildcard expansion depends on
|
||||
// whether the target peer supports IPv6. Routes and firewall rules are already
|
||||
// filtered at the source (network map builder).
|
||||
@@ -162,12 +162,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),
|
||||
}
|
||||
|
||||
@@ -276,7 +276,7 @@ func TestToNetbirdConfig_RelayInvariant(t *testing.T) {
|
||||
settings := &types.Settings{MetricsPushEnabled: true}
|
||||
|
||||
t.Run("nil server config returns nil config", func(t *testing.T) {
|
||||
nbCfg := toNetbirdConfig(nil, nil, nil, nil, settings)
|
||||
nbCfg := toNetbirdConfig(nil, nil, nil, nil, types.TwinAccountSettings(settings))
|
||||
assert.Nil(t, nbCfg, "fan-out updates must not carry a partial NetbirdConfig even when settings are present")
|
||||
})
|
||||
|
||||
@@ -291,7 +291,7 @@ func TestToNetbirdConfig_RelayInvariant(t *testing.T) {
|
||||
}
|
||||
relayToken := &Token{Payload: "token-payload", Signature: "token-signature"}
|
||||
|
||||
nbCfg := toNetbirdConfig(cfg, nil, relayToken, nil, settings)
|
||||
nbCfg := toNetbirdConfig(cfg, nil, relayToken, nil, types.TwinAccountSettings(settings))
|
||||
require.NotNil(t, nbCfg)
|
||||
require.NotNil(t, nbCfg.Relay, "non-nil NetbirdConfig must include the relay section")
|
||||
assert.Equal(t, cfg.Relay.Addresses, nbCfg.Relay.Urls, "relay URLs should match the server config")
|
||||
|
||||
@@ -920,8 +920,8 @@ 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, network, s.networkMapController.GetDNSDomain(settings), settings, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, enableSSH),
|
||||
NetbirdConfig: toNetbirdConfig(s.config, nil, relayToken, nil, types.TwinAccountSettings(settings)),
|
||||
PeerConfig: toPeerConfig(types.TwinPeer(peer), types.TwinNetwork(network), s.networkMapController.GetDNSDomain(settings), types.TwinAccountSettings(settings), s.config.HttpConfig, s.config.DeviceAuthorizationFlow, enableSSH, false),
|
||||
Checks: toProtocolChecks(ctx, postureChecks),
|
||||
}
|
||||
|
||||
@@ -1052,9 +1052,9 @@ func (s *Server) sendInitialSync(ctx context.Context, peerKey wgtypes.Key, peer
|
||||
log.WithContext(ctx).Errorf("failed to build components for peer %s on initial sync: %v", peer.ID, err)
|
||||
return status.Errorf(codes.Internal, "failed to build initial sync envelope")
|
||||
}
|
||||
plainResp = ToComponentSyncResponse(ctx, s.config, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, freshPeer, turnToken, relayToken, components, proxyPatch, dnsName, freshPostureChecks, settings, settings.Extra, peerGroups, freshDnsFwdPort)
|
||||
plainResp = ToComponentSyncResponse(ctx, s.config, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, types.TwinPeer(freshPeer), turnToken, relayToken, components, proxyPatch, dnsName, freshPostureChecks, types.TwinAccountSettings(settings), settings.Extra, peerGroups, freshDnsFwdPort)
|
||||
} else {
|
||||
plainResp = ToSyncResponse(ctx, s.config, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, peer, turnToken, relayToken, networkMap, dnsName, postureChecks, nil, settings, settings.Extra, peerGroups, dnsFwdPort)
|
||||
plainResp = ToSyncResponse(ctx, s.config, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, types.TwinPeer(peer), turnToken, relayToken, networkMap, dnsName, postureChecks, nil, types.TwinAccountSettings(settings), settings.Extra, peerGroups, dnsFwdPort)
|
||||
}
|
||||
|
||||
key, err := s.secretsManager.GetWGKey()
|
||||
|
||||
@@ -3330,7 +3330,7 @@ func createManager(t testing.TB) (*DefaultAccountManager, *update_channel.PeersU
|
||||
|
||||
updateManager := update_channel.NewPeersUpdateManager(metrics)
|
||||
requestBuffer := NewAccountRequestBuffer(ctx, store)
|
||||
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{})
|
||||
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{}, nil)
|
||||
manager, err := BuildManager(ctx, &config.Config{}, store, networkMapController, job.NewJobManager(nil, store, peersManager), nil, "", eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
|
||||
@@ -234,7 +234,7 @@ func createDNSManager(t *testing.T) (*DefaultAccountManager, error) {
|
||||
|
||||
updateManager := update_channel.NewPeersUpdateManager(metrics)
|
||||
requestBuffer := NewAccountRequestBuffer(ctx, store)
|
||||
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.test", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{})
|
||||
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.test", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{}, nil)
|
||||
|
||||
return BuildManager(context.Background(), nil, store, networkMapController, job.NewJobManager(nil, store, peersManager), nil, "", eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
|
||||
}
|
||||
|
||||
@@ -6,7 +6,6 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/account"
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
|
||||
"github.com/netbirdio/netbird/management/server/permissions"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/modules"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/operations"
|
||||
@@ -31,10 +30,6 @@ type managerImpl struct {
|
||||
accountManager account.Manager
|
||||
}
|
||||
|
||||
func eventMetaResource(group *types.Group, resource *resourceTypes.NetworkResource) map[string]any {
|
||||
return map[string]any{"name": group.Name, "id": group.ID, "resource_name": resource.Name, "resource_id": resource.ID, "resource_type": resource.Type}
|
||||
}
|
||||
|
||||
type mockManager struct {
|
||||
}
|
||||
|
||||
@@ -114,7 +109,7 @@ func (m *managerImpl) AddResourceToGroupInTransaction(ctx context.Context, trans
|
||||
}
|
||||
|
||||
event := func() {
|
||||
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceAddedToGroup, eventMetaResource(group, networkResource))
|
||||
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceAddedToGroup, group.EventMetaResource(networkResource))
|
||||
}
|
||||
|
||||
return event, nil
|
||||
@@ -138,7 +133,7 @@ func (m *managerImpl) RemoveResourceFromGroupInTransaction(ctx context.Context,
|
||||
}
|
||||
|
||||
event := func() {
|
||||
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceRemovedFromGroup, eventMetaResource(group, networkResource))
|
||||
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceRemovedFromGroup, group.EventMetaResource(networkResource))
|
||||
}
|
||||
|
||||
return event, nil
|
||||
|
||||
@@ -446,7 +446,7 @@ func (h *Handler) GetAccessiblePeers(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
netMap := account.GetPeerNetworkMapFromComponents(ctx, peerID, dns.CustomZone{}, nil, validPeers, account.GetResourcePoliciesMap(), account.GetResourceRoutersMap(), nil, account.GetActiveGroupUsers())
|
||||
|
||||
util.WriteJSONObject(ctx, w, toAccessiblePeers(netMap, account.Peers, dnsDomain))
|
||||
util.WriteJSONObject(ctx, w, toAccessiblePeers(account.Peers, netMap, dnsDomain))
|
||||
}
|
||||
|
||||
func (h *Handler) CreateTemporaryAccess(w http.ResponseWriter, r *http.Request) {
|
||||
@@ -534,20 +534,22 @@ func (h *Handler) CreateTemporaryAccess(w http.ResponseWriter, r *http.Request)
|
||||
util.WriteJSONObject(r.Context(), w, resp)
|
||||
}
|
||||
|
||||
// toAccessiblePeers rehydrates the calculated map's component peers into the
|
||||
// account's full peer objects, which carry the location/status/meta fields
|
||||
// the API response needs.
|
||||
func toAccessiblePeers(netMap *types.NetworkMap, accountPeers map[string]*nbpeer.Peer, dnsDomain string) []api.AccessiblePeer {
|
||||
// toAccessiblePeers resolves the twin peers in netMap back to the full account
|
||||
// peers (by ID) so the API response keeps Status/Name/OS/GeoNameID, which the
|
||||
// slim netmap twins intentionally don't carry.
|
||||
func toAccessiblePeers(accountPeers map[string]*nbpeer.Peer, netMap *types.NetworkMap, dnsDomain string) []api.AccessiblePeer {
|
||||
accessiblePeers := make([]api.AccessiblePeer, 0, len(netMap.Peers)+len(netMap.OfflinePeers))
|
||||
add := func(peers []*types.ComponentPeer) {
|
||||
for _, p := range peers {
|
||||
if peer := accountPeers[p.ID]; peer != nil {
|
||||
accessiblePeers = append(accessiblePeers, peerToAccessiblePeer(peer, dnsDomain))
|
||||
}
|
||||
appendByID := func(id string) {
|
||||
if p, ok := accountPeers[id]; ok && p != nil {
|
||||
accessiblePeers = append(accessiblePeers, peerToAccessiblePeer(p, dnsDomain))
|
||||
}
|
||||
}
|
||||
add(netMap.Peers)
|
||||
add(netMap.OfflinePeers)
|
||||
for _, p := range netMap.Peers {
|
||||
appendByID(p.ID)
|
||||
}
|
||||
for _, p := range netMap.OfflinePeers {
|
||||
appendByID(p.ID)
|
||||
}
|
||||
|
||||
return accessiblePeers
|
||||
}
|
||||
|
||||
@@ -96,7 +96,7 @@ func BuildApiBlackBoxWithDBState(t testing_tools.TB, sqlFile string, expectedPee
|
||||
}
|
||||
|
||||
requestBuffer := server.NewAccountRequestBuffer(ctx, store)
|
||||
networkMapController := controller.NewController(ctx, store, metrics, peersUpdateManager, requestBuffer, server.MockIntegratedValidator{}, settingsManager, "", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peersManager), &config.Config{})
|
||||
networkMapController := controller.NewController(ctx, store, metrics, peersUpdateManager, requestBuffer, server.MockIntegratedValidator{}, settingsManager, "", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peersManager), &config.Config{}, nil)
|
||||
am, err := server.BuildManager(ctx, nil, store, networkMapController, jobManager, nil, "", &activity.InMemoryEventStore{}, geoMock, false, validatorMock, metrics, proxyController, settingsManager, permissionsManager, false, cacheStore)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create manager: %v", err)
|
||||
@@ -226,7 +226,7 @@ func BuildApiBlackBoxWithDBStateAndPeerChannel(t testing_tools.TB, sqlFile strin
|
||||
}
|
||||
|
||||
requestBuffer := server.NewAccountRequestBuffer(ctx, store)
|
||||
networkMapController := controller.NewController(ctx, store, metrics, peersUpdateManager, requestBuffer, server.MockIntegratedValidator{}, settingsManager, "", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peersManager), &config.Config{})
|
||||
networkMapController := controller.NewController(ctx, store, metrics, peersUpdateManager, requestBuffer, server.MockIntegratedValidator{}, settingsManager, "", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peersManager), &config.Config{}, nil)
|
||||
am, err := server.BuildManager(ctx, nil, store, networkMapController, jobManager, nil, "", &activity.InMemoryEventStore{}, geoMock, false, validatorMock, metrics, proxyController, settingsManager, permissionsManager, false, cacheStore)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create manager: %v", err)
|
||||
|
||||
@@ -92,7 +92,7 @@ func createManagerWithEmbeddedIdP(t testing.TB) (*DefaultAccountManager, *update
|
||||
|
||||
updateManager := update_channel.NewPeersUpdateManager(metrics)
|
||||
requestBuffer := NewAccountRequestBuffer(ctx, testStore)
|
||||
networkMapController := controller.NewController(ctx, testStore, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(testStore, peersManager), &config.Config{})
|
||||
networkMapController := controller.NewController(ctx, testStore, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(testStore, peersManager), &config.Config{}, nil)
|
||||
manager, err := BuildManager(ctx, &config.Config{}, testStore, networkMapController, job.NewJobManager(nil, testStore, peersManager), idpManager, "", eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
// UpdateIntegratedValidator updates the integrated validator groups for a specified account.
|
||||
@@ -109,7 +110,7 @@ func (am *DefaultAccountManager) GetValidatedPeers(ctx context.Context, accountI
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
validPeers, err := am.integratedPeerValidator.GetValidatedPeers(ctx, accountID, groups, peers, settings.Extra)
|
||||
validPeers, err := am.integratedPeerValidator.GetValidatedPeers(ctx, accountID, types.TwinGroups(groups), types.TwinPeers(peers), settings.Extra)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
@@ -138,7 +139,7 @@ func (a MockIntegratedValidator) ValidatePeer(_ context.Context, update *nbpeer.
|
||||
return update, false, nil
|
||||
}
|
||||
|
||||
func (a MockIntegratedValidator) GetValidatedPeers(_ context.Context, accountID string, groups []*types.Group, peers []*nbpeer.Peer, extraSettings *types.ExtraSettings) (map[string]struct{}, error) {
|
||||
func (a MockIntegratedValidator) GetValidatedPeers(_ context.Context, accountID string, groups []*nmdata.Group, peers []*nmdata.Peer, extraSettings *types.ExtraSettings) (map[string]struct{}, error) {
|
||||
validatedPeers := make(map[string]struct{})
|
||||
for _, peer := range peers {
|
||||
validatedPeers[peer.ID] = struct{}{}
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
@@ -14,7 +15,7 @@ type IntegratedValidator interface {
|
||||
ValidatePeer(ctx context.Context, update *nbpeer.Peer, peer *nbpeer.Peer, userID string, accountID string, dnsDomain string, peersGroup []string, extraSettings *types.ExtraSettings) (*nbpeer.Peer, bool, error)
|
||||
PreparePeer(ctx context.Context, accountID string, peer *nbpeer.Peer, peersGroup []string, extraSettings *types.ExtraSettings, temporary bool) *nbpeer.Peer
|
||||
IsNotValidPeer(ctx context.Context, accountID string, peer *nbpeer.Peer, peersGroup []string, extraSettings *types.ExtraSettings) (bool, bool, error)
|
||||
GetValidatedPeers(ctx context.Context, accountID string, groups []*types.Group, peers []*nbpeer.Peer, extraSettings *types.ExtraSettings) (map[string]struct{}, error)
|
||||
GetValidatedPeers(ctx context.Context, accountID string, groups []*nmdata.Group, peers []*nmdata.Peer, extraSettings *types.ExtraSettings) (map[string]struct{}, error)
|
||||
GetInvalidPeers(ctx context.Context, accountID string, extraSettings *types.ExtraSettings) (map[string]string, error)
|
||||
PeerDeleted(ctx context.Context, accountID, peerID string, extraSettings *types.ExtraSettings) error
|
||||
SetPeerInvalidationListener(fn func(accountID string, peerIDs []string))
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/settings"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
@@ -35,7 +36,7 @@ func (v *IntegratedValidatorImpl) IsNotValidPeer(_ context.Context, _ string, _
|
||||
return false, false, nil
|
||||
}
|
||||
|
||||
func (v *IntegratedValidatorImpl) GetValidatedPeers(_ context.Context, _ string, _ []*types.Group, peers []*nbpeer.Peer, _ *types.ExtraSettings) (map[string]struct{}, error) {
|
||||
func (v *IntegratedValidatorImpl) GetValidatedPeers(_ context.Context, _ string, _ []*nmdata.Group, peers []*nmdata.Peer, _ *types.ExtraSettings) (map[string]struct{}, error) {
|
||||
validatedPeers := make(map[string]struct{})
|
||||
for _, p := range peers {
|
||||
validatedPeers[p.ID] = struct{}{}
|
||||
|
||||
@@ -376,7 +376,7 @@ func startManagementForTest(t *testing.T, testFile string, config *config.Config
|
||||
return nil, nil, "", cleanup, err
|
||||
}
|
||||
|
||||
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeralMgr, config)
|
||||
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeralMgr, config, nil)
|
||||
accountManager, err := BuildManager(ctx, nil, store, networkMapController, jobManager, nil, "",
|
||||
eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
|
||||
|
||||
|
||||
@@ -216,7 +216,7 @@ func startServer(
|
||||
|
||||
updateManager := update_channel.NewPeersUpdateManager(metrics)
|
||||
requestBuffer := server.NewAccountRequestBuffer(ctx, str)
|
||||
networkMapController := controller.NewController(ctx, str, metrics, updateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(str, peers.NewManager(str, permissionsManager)), config)
|
||||
networkMapController := controller.NewController(ctx, str, metrics, updateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(str, peers.NewManager(str, permissionsManager)), config, nil)
|
||||
|
||||
accountManager, err := server.BuildManager(
|
||||
context.Background(),
|
||||
|
||||
@@ -803,7 +803,7 @@ func createNSManager(t *testing.T) (*DefaultAccountManager, error) {
|
||||
|
||||
updateManager := update_channel.NewPeersUpdateManager(metrics)
|
||||
requestBuffer := NewAccountRequestBuffer(ctx, store)
|
||||
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{})
|
||||
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{}, nil)
|
||||
|
||||
return BuildManager(context.Background(), nil, store, networkMapController, job.NewJobManager(nil, store, peersManager), nil, "", eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
|
||||
}
|
||||
|
||||
@@ -14,7 +14,6 @@ import (
|
||||
nbDomain "github.com/netbirdio/netbird/shared/management/domain"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
sharedTypes "github.com/netbirdio/netbird/shared/management/types"
|
||||
)
|
||||
|
||||
type NetworkResourceType string
|
||||
@@ -65,27 +64,6 @@ func NewNetworkResource(accountID, networkID, name, description, address string,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ToComponent converts the resource to its self-contained components
|
||||
// representation. Returns nil for a nil resource.
|
||||
func (n *NetworkResource) ToComponent() *sharedTypes.ComponentResource {
|
||||
if n == nil {
|
||||
return nil
|
||||
}
|
||||
return &sharedTypes.ComponentResource{
|
||||
ID: n.ID,
|
||||
PublicID: n.PublicID,
|
||||
NetworkID: n.NetworkID,
|
||||
AccountID: n.AccountID,
|
||||
Name: n.Name,
|
||||
Description: n.Description,
|
||||
Type: sharedTypes.ComponentResourceType(n.Type),
|
||||
Address: n.Address,
|
||||
Domain: n.Domain,
|
||||
Prefix: n.Prefix,
|
||||
Enabled: n.Enabled,
|
||||
}
|
||||
}
|
||||
|
||||
func (n *NetworkResource) ToAPIResponse(groups []api.GroupMinimum) *api.NetworkResource {
|
||||
addr := n.Prefix.String()
|
||||
if n.Type == Domain {
|
||||
|
||||
@@ -7,7 +7,6 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/networks/types"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
sharedTypes "github.com/netbirdio/netbird/shared/management/types"
|
||||
)
|
||||
|
||||
type NetworkRouter struct {
|
||||
@@ -22,36 +21,6 @@ type NetworkRouter struct {
|
||||
Enabled bool
|
||||
}
|
||||
|
||||
// ToComponent converts the router to its self-contained components
|
||||
// representation. Returns nil for a nil router.
|
||||
func (n *NetworkRouter) ToComponent() *sharedTypes.ComponentRouter {
|
||||
if n == nil {
|
||||
return nil
|
||||
}
|
||||
return &sharedTypes.ComponentRouter{
|
||||
NetworkID: n.NetworkID,
|
||||
PublicID: n.PublicID,
|
||||
Peer: n.Peer,
|
||||
PeerGroups: n.PeerGroups,
|
||||
Masquerade: n.Masquerade,
|
||||
Metric: n.Metric,
|
||||
Enabled: n.Enabled,
|
||||
}
|
||||
}
|
||||
|
||||
// ToComponentMap converts a peer-keyed router map to its components
|
||||
// representation.
|
||||
func ToComponentMap(routers map[string]*NetworkRouter) map[string]*sharedTypes.ComponentRouter {
|
||||
if routers == nil {
|
||||
return nil
|
||||
}
|
||||
out := make(map[string]*sharedTypes.ComponentRouter, len(routers))
|
||||
for id, r := range routers {
|
||||
out[id] = r.ToComponent()
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func NewNetworkRouter(accountID string, networkID string, peer string, peerGroups []string, masquerade bool, metric int, enabled bool) (*NetworkRouter, error) {
|
||||
r := &NetworkRouter{
|
||||
ID: xid.New().String(),
|
||||
|
||||
@@ -21,6 +21,7 @@ import (
|
||||
"github.com/netbirdio/netbird/management/server/permissions/modules"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/operations"
|
||||
"github.com/netbirdio/netbird/shared/management/domain"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/posture"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
@@ -405,7 +406,7 @@ func (am *DefaultAccountManager) CreatePeerJob(ctx context.Context, accountID, p
|
||||
return status.NewPeerNotPartOfAccountError()
|
||||
}
|
||||
|
||||
meetMinVer, err := version.MeetsMinVersion(remoteJobsMinVer, p.Meta.WtVersion)
|
||||
meetMinVer, err := posture.MeetsMinVersion(remoteJobsMinVer, p.Meta.WtVersion)
|
||||
if !version.IsDevelopmentVersion(p.Meta.WtVersion) && (!meetMinVer || err != nil) {
|
||||
return status.Errorf(status.PreconditionFailed, "peer version %s does not meet the minimum required version %s for remote jobs", p.Meta.WtVersion, remoteJobsMinVer)
|
||||
}
|
||||
@@ -1588,7 +1589,7 @@ func affectedPeerIDsFromNetworkMap(nmap *types.NetworkMap, selfPeerID string) []
|
||||
}
|
||||
seen := make(map[string]struct{}, len(nmap.Peers)+len(nmap.OfflinePeers))
|
||||
ids := make([]string, 0, len(nmap.Peers)+len(nmap.OfflinePeers))
|
||||
add := func(peers []*types.ComponentPeer) {
|
||||
add := func(peers []*nmdata.Peer) {
|
||||
for _, p := range peers {
|
||||
if p == nil || p.ID == "" || p.ID == selfPeerID {
|
||||
continue
|
||||
|
||||
@@ -13,7 +13,6 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/util"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
sharedTypes "github.com/netbirdio/netbird/shared/management/types"
|
||||
)
|
||||
|
||||
// Peer capability constants mirror the proto enum values.
|
||||
@@ -206,35 +205,6 @@ func (p *Peer) AddedWithSSOLogin() bool {
|
||||
return p.UserID != ""
|
||||
}
|
||||
|
||||
// ToComponent converts the peer to its self-contained components
|
||||
// representation, carrying exactly the subset of peer data that crosses the
|
||||
// components wire format. Returns nil for a nil peer so callers can convert
|
||||
// possibly-missing peers without guarding.
|
||||
func (p *Peer) ToComponent() *sharedTypes.ComponentPeer {
|
||||
if p == nil {
|
||||
return nil
|
||||
}
|
||||
cp := &sharedTypes.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
|
||||
}
|
||||
|
||||
// HasCapability reports whether the peer has the given capability.
|
||||
func (p *Peer) HasCapability(capability int32) bool {
|
||||
return slices.Contains(p.Meta.Capabilities, capability)
|
||||
|
||||
@@ -57,6 +57,7 @@ import (
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
nbroute "github.com/netbirdio/netbird/route"
|
||||
"github.com/netbirdio/netbird/shared/management/domain"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
@@ -1091,22 +1092,22 @@ func TestToSyncResponse(t *testing.T) {
|
||||
Signature: "turn-pass",
|
||||
}
|
||||
networkMap := &types.NetworkMap{
|
||||
Network: &types.Network{Net: *ipnet, Serial: 1000},
|
||||
Peers: []*types.ComponentPeer{{
|
||||
Network: &nmdata.Network{Net: *ipnet, Serial: 1000},
|
||||
Peers: []*nmdata.Peer{{
|
||||
IP: netip.MustParseAddr("192.168.1.2"),
|
||||
IPv6: netip.MustParseAddr("fd00::2"),
|
||||
Key: "peer2-key",
|
||||
DNSLabel: "peer2",
|
||||
SSHEnabled: true,
|
||||
SSHKey: "peer2-ssh-key"}},
|
||||
OfflinePeers: []*types.ComponentPeer{{
|
||||
OfflinePeers: []*nmdata.Peer{{
|
||||
IP: netip.MustParseAddr("192.168.1.3"),
|
||||
IPv6: netip.MustParseAddr("fd00::3"),
|
||||
Key: "peer3-key",
|
||||
DNSLabel: "peer3",
|
||||
SSHEnabled: true,
|
||||
SSHKey: "peer3-ssh-key"}},
|
||||
Routes: []*nbroute.Route{
|
||||
Routes: []*nmdata.Route{
|
||||
{
|
||||
ID: "route1",
|
||||
Network: netip.MustParsePrefix("10.0.0.0/24"),
|
||||
@@ -1180,7 +1181,7 @@ func TestToSyncResponse(t *testing.T) {
|
||||
}
|
||||
dnsCache := &cache.DNSConfigCache{}
|
||||
accountSettings := &types.Settings{RoutingPeerDNSResolutionEnabled: true}
|
||||
response := grpc.ToSyncResponse(context.Background(), config, config.HttpConfig, config.DeviceAuthorizationFlow, peer, turnRelayToken, turnRelayToken, networkMap, dnsName, checks, dnsCache, accountSettings, nil, []string{}, int64(dnsForwarderPort))
|
||||
response := grpc.ToSyncResponse(context.Background(), config, config.HttpConfig, config.DeviceAuthorizationFlow, types.TwinPeer(peer), turnRelayToken, turnRelayToken, networkMap, dnsName, checks, dnsCache, types.TwinAccountSettings(accountSettings), nil, []string{}, int64(dnsForwarderPort))
|
||||
|
||||
assert.NotNil(t, response)
|
||||
// assert peer config
|
||||
@@ -1300,7 +1301,7 @@ func Test_RegisterPeerByUser(t *testing.T) {
|
||||
|
||||
updateManager := update_channel.NewPeersUpdateManager(metrics)
|
||||
requestBuffer := NewAccountRequestBuffer(ctx, s)
|
||||
networkMapController := controller.NewController(ctx, s, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(s, peers.NewManager(s, permissionsManager)), &config.Config{})
|
||||
networkMapController := controller.NewController(ctx, s, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(s, peers.NewManager(s, permissionsManager)), &config.Config{}, nil)
|
||||
|
||||
am, err := BuildManager(context.Background(), nil, s, networkMapController, job.NewJobManager(nil, s, peersManager), nil, "", eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
|
||||
assert.NoError(t, err)
|
||||
@@ -1391,7 +1392,7 @@ func Test_RegisterPeerBySetupKey(t *testing.T) {
|
||||
|
||||
updateManager := update_channel.NewPeersUpdateManager(metrics)
|
||||
requestBuffer := NewAccountRequestBuffer(ctx, s)
|
||||
networkMapController := controller.NewController(ctx, s, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(s, peers.NewManager(s, permissionsManager)), &config.Config{})
|
||||
networkMapController := controller.NewController(ctx, s, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(s, peers.NewManager(s, permissionsManager)), &config.Config{}, nil)
|
||||
|
||||
am, err := BuildManager(context.Background(), nil, s, networkMapController, job.NewJobManager(nil, s, peersManager), nil, "", eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
|
||||
assert.NoError(t, err)
|
||||
@@ -1550,7 +1551,7 @@ func Test_RegisterPeerRollbackOnFailure(t *testing.T) {
|
||||
|
||||
updateManager := update_channel.NewPeersUpdateManager(metrics)
|
||||
requestBuffer := NewAccountRequestBuffer(ctx, s)
|
||||
networkMapController := controller.NewController(ctx, s, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(s, peers.NewManager(s, permissionsManager)), &config.Config{})
|
||||
networkMapController := controller.NewController(ctx, s, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(s, peers.NewManager(s, permissionsManager)), &config.Config{}, nil)
|
||||
|
||||
am, err := BuildManager(context.Background(), nil, s, networkMapController, job.NewJobManager(nil, s, peersManager), nil, "", eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
|
||||
assert.NoError(t, err)
|
||||
@@ -1635,7 +1636,7 @@ func Test_LoginPeer(t *testing.T) {
|
||||
|
||||
updateManager := update_channel.NewPeersUpdateManager(metrics)
|
||||
requestBuffer := NewAccountRequestBuffer(ctx, s)
|
||||
networkMapController := controller.NewController(ctx, s, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(s, peers.NewManager(s, permissionsManager)), &config.Config{})
|
||||
networkMapController := controller.NewController(ctx, s, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(s, peers.NewManager(s, permissionsManager)), &config.Config{}, nil)
|
||||
|
||||
am, err := BuildManager(context.Background(), nil, s, networkMapController, job.NewJobManager(nil, s, peersManager), nil, "", eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
|
||||
assert.NoError(t, err)
|
||||
|
||||
@@ -3,9 +3,11 @@ package posture
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/hashicorp/go-version"
|
||||
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
nbversion "github.com/netbirdio/netbird/version"
|
||||
)
|
||||
|
||||
type NBVersionCheck struct {
|
||||
@@ -14,8 +16,14 @@ type NBVersionCheck struct {
|
||||
|
||||
var _ Check = (*NBVersionCheck)(nil)
|
||||
|
||||
// sanitizeVersion removes anything after the pre-release tag (e.g., "-dev", "-alpha", etc.)
|
||||
func sanitizeVersion(version string) string {
|
||||
parts := strings.Split(version, "-")
|
||||
return parts[0]
|
||||
}
|
||||
|
||||
func (n *NBVersionCheck) Check(ctx context.Context, peer nbpeer.Peer) (bool, error) {
|
||||
meetsMin, err := nbversion.MeetsMinVersion(n.MinVersion, peer.Meta.WtVersion)
|
||||
meetsMin, err := MeetsMinVersion(n.MinVersion, peer.Meta.WtVersion)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
@@ -40,3 +48,21 @@ func (n *NBVersionCheck) Validate() error {
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// MeetsMinVersion checks if the peer's version meets or exceeds the minimum required version
|
||||
func MeetsMinVersion(minVer, peerVer string) (bool, error) {
|
||||
peerVer = sanitizeVersion(peerVer)
|
||||
minVer = sanitizeVersion(minVer)
|
||||
|
||||
peerNBVer, err := version.NewVersion(peerVer)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
constraints, err := version.NewConstraint(">= " + minVer)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
return constraints.Check(peerNBVer), nil
|
||||
}
|
||||
|
||||
@@ -139,3 +139,68 @@ func TestNBVersionCheck_Validate(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMeetsMinVersion(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
minVer string
|
||||
peerVer string
|
||||
want bool
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "Peer version greater than min version",
|
||||
minVer: "0.26.0",
|
||||
peerVer: "0.60.1",
|
||||
want: true,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "Peer version equals min version",
|
||||
minVer: "1.0.0",
|
||||
peerVer: "1.0.0",
|
||||
want: true,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "Peer version less than min version",
|
||||
minVer: "1.0.0",
|
||||
peerVer: "0.9.9",
|
||||
want: false,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "Peer version with pre-release tag greater than min version",
|
||||
minVer: "1.0.0",
|
||||
peerVer: "1.0.1-alpha",
|
||||
want: true,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "Invalid peer version format",
|
||||
minVer: "1.0.0",
|
||||
peerVer: "dev",
|
||||
want: false,
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "Invalid min version format",
|
||||
minVer: "invalid.version",
|
||||
peerVer: "1.0.0",
|
||||
want: false,
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := MeetsMinVersion(tt.minVer, tt.peerVer)
|
||||
if tt.wantErr {
|
||||
assert.Error(t, err)
|
||||
} else {
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
assert.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1201,7 +1201,7 @@ func TestGetNetworkMap_RouteSync(t *testing.T) {
|
||||
peer1Routes, err := am.GetNetworkMap(context.Background(), peer1ID)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, peer1Routes.Routes, 1, "we should receive one route for peer1")
|
||||
require.True(t, expectedRoute.Equal(peer1Routes.Routes[0]), "received route should be equal")
|
||||
require.True(t, types.TwinRoute(expectedRoute).Equal(peer1Routes.Routes[0]), "received route should be equal")
|
||||
|
||||
peer2Routes, err := am.GetNetworkMap(context.Background(), peer2ID)
|
||||
require.NoError(t, err)
|
||||
@@ -1299,7 +1299,7 @@ func createRouterManager(t *testing.T) (*DefaultAccountManager, *update_channel.
|
||||
|
||||
updateManager := update_channel.NewPeersUpdateManager(metrics)
|
||||
requestBuffer := NewAccountRequestBuffer(ctx, store)
|
||||
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{})
|
||||
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{}, nil)
|
||||
|
||||
am, err := BuildManager(context.Background(), nil, store, networkMapController, job.NewJobManager(nil, store, peersManager), nil, "", eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
|
||||
if err != nil {
|
||||
|
||||
@@ -9,7 +9,6 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/hashicorp/go-multierror"
|
||||
"github.com/miekg/dns"
|
||||
"github.com/rs/xid"
|
||||
log "github.com/sirupsen/logrus"
|
||||
@@ -18,8 +17,6 @@ import (
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
proxydomain "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
|
||||
"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"
|
||||
networkTypes "github.com/netbirdio/netbird/management/server/networks/types"
|
||||
@@ -28,6 +25,8 @@ import (
|
||||
"github.com/netbirdio/netbird/management/server/util"
|
||||
"github.com/netbirdio/netbird/route"
|
||||
"github.com/netbirdio/netbird/shared/management/domain"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
@@ -382,94 +381,11 @@ func peerInDistributionGroups(peerGroups LookupMap, distributionGroups []string)
|
||||
}
|
||||
|
||||
func (a *Account) GetPeersCustomZone(ctx context.Context, dnsDomain string) nbdns.CustomZone {
|
||||
var merr *multierror.Error
|
||||
|
||||
if dnsDomain == "" {
|
||||
log.WithContext(ctx).Error("no dns domain is set, returning empty zone")
|
||||
return nbdns.CustomZone{}
|
||||
twins := make(map[string]*nmdata.Peer, len(a.Peers))
|
||||
for id, p := range a.Peers {
|
||||
twins[id] = twinPeer(p)
|
||||
}
|
||||
|
||||
customZone := nbdns.CustomZone{
|
||||
Domain: dns.Fqdn(dnsDomain),
|
||||
Records: make([]nbdns.SimpleRecord, 0, len(a.Peers)),
|
||||
}
|
||||
|
||||
domainSuffix := "." + dnsDomain
|
||||
|
||||
ipv6AllowedPeers := a.peerIPv6AllowedSet()
|
||||
|
||||
var sb strings.Builder
|
||||
for _, peer := range a.Peers {
|
||||
if peer.DNSLabel == "" {
|
||||
merr = multierror.Append(merr, fmt.Errorf("peer %s has an empty DNS label", peer.Name))
|
||||
continue
|
||||
}
|
||||
|
||||
sb.Grow(len(peer.DNSLabel) + len(domainSuffix))
|
||||
sb.WriteString(peer.DNSLabel)
|
||||
sb.WriteString(domainSuffix)
|
||||
|
||||
fqdn := sb.String()
|
||||
customZone.Records = append(customZone.Records, nbdns.SimpleRecord{
|
||||
Name: fqdn,
|
||||
Type: int(dns.TypeA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: defaultTTL,
|
||||
RData: peer.IP.String(),
|
||||
})
|
||||
// Only advertise AAAA for peers that have a valid IPv6, whose client supports it,
|
||||
// and that belong to an IPv6-enabled group. Old clients don't configure v6 on their
|
||||
// WireGuard interface, so resolving their AAAA causes connections to hang.
|
||||
// Capability changes (client upgrade/downgrade, --disable-ipv6 toggle) propagate
|
||||
// to other peers via SyncPeer/LoginPeer regardless of version change, so AAAA
|
||||
// records refresh when a peer first reports the IPv6 overlay capability.
|
||||
_, peerAllowed := ipv6AllowedPeers[peer.ID]
|
||||
hasIPv6 := peer.IPv6.IsValid() && peer.SupportsIPv6() && peerAllowed
|
||||
if hasIPv6 {
|
||||
customZone.Records = append(customZone.Records, nbdns.SimpleRecord{
|
||||
Name: fqdn,
|
||||
Type: int(dns.TypeAAAA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: defaultTTL,
|
||||
RData: peer.IPv6.String(),
|
||||
})
|
||||
}
|
||||
sb.Reset()
|
||||
|
||||
for _, extraLabel := range peer.ExtraDNSLabels {
|
||||
sb.Grow(len(extraLabel) + len(domainSuffix))
|
||||
sb.WriteString(extraLabel)
|
||||
sb.WriteString(domainSuffix)
|
||||
|
||||
extraFqdn := sb.String()
|
||||
customZone.Records = append(customZone.Records, nbdns.SimpleRecord{
|
||||
Name: extraFqdn,
|
||||
Type: int(dns.TypeA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: defaultTTL,
|
||||
RData: peer.IP.String(),
|
||||
})
|
||||
if hasIPv6 {
|
||||
customZone.Records = append(customZone.Records, nbdns.SimpleRecord{
|
||||
Name: extraFqdn,
|
||||
Type: int(dns.TypeAAAA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: defaultTTL,
|
||||
RData: peer.IPv6.String(),
|
||||
})
|
||||
}
|
||||
sb.Reset()
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
go func() {
|
||||
if merr != nil {
|
||||
log.WithContext(ctx).Errorf("error generating custom zone for account %s: %v", a.Id, merr)
|
||||
}
|
||||
}()
|
||||
|
||||
return customZone
|
||||
return fromTwinCustomZone(networkmap.PeersCustomZone(ctx, a.Id, dnsDomain, twins, a.peerIPv6AllowedSet()))
|
||||
}
|
||||
|
||||
// GetExpiredPeers returns peers that have been expired
|
||||
@@ -1062,6 +978,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 {
|
||||
@@ -1082,7 +1046,6 @@ func (a *Account) connResourcesGenerator(ctx context.Context, targetPeer *nbpeer
|
||||
peersExists := make(map[string]struct{})
|
||||
rules := make([]*FirewallRule, 0)
|
||||
peers := make([]*nbpeer.Peer, 0)
|
||||
targetComponent := targetPeer.ToComponent()
|
||||
|
||||
return func(rule *PolicyRule, groupPeers []*nbpeer.Peer, direction int) {
|
||||
for _, peer := range groupPeers {
|
||||
@@ -1118,10 +1081,10 @@ func (a *Account) connResourcesGenerator(ctx context.Context, targetPeer *nbpeer
|
||||
if len(rule.Ports) == 0 && len(rule.PortRanges) == 0 {
|
||||
rules = append(rules, &fr)
|
||||
} else {
|
||||
rules = append(rules, ExpandPortsAndRanges(fr, rule, targetComponent)...)
|
||||
rules = append(rules, ExpandPortsAndRanges(fr, rule, targetPeer)...)
|
||||
}
|
||||
|
||||
rules = AppendIPv6FirewallRule(rules, rulesExists, peer.ToComponent(), targetComponent, rule, FirewallRuleContext{
|
||||
rules = AppendIPv6FirewallRule(rules, rulesExists, peer, targetPeer, rule, FirewallRuleContext{
|
||||
Direction: direction,
|
||||
DirStr: strconv.Itoa(direction),
|
||||
ProtocolStr: string(protocol),
|
||||
@@ -1281,7 +1244,7 @@ func (a *Account) getRouteFirewallRules(ctx context.Context, peerID string, poli
|
||||
return fwRules
|
||||
}
|
||||
|
||||
func (a *Account) getRulePeers(rule *PolicyRule, postureChecks []string, peerID string, distributionPeers map[string]struct{}, validatedPeersMap map[string]struct{}) []*ComponentPeer {
|
||||
func (a *Account) getRulePeers(rule *PolicyRule, postureChecks []string, peerID string, distributionPeers map[string]struct{}, validatedPeersMap map[string]struct{}) []*nbpeer.Peer {
|
||||
distPeersWithPolicy := make(map[string]struct{})
|
||||
for _, id := range rule.Sources {
|
||||
group := a.Groups[id]
|
||||
@@ -1308,13 +1271,13 @@ func (a *Account) getRulePeers(rule *PolicyRule, postureChecks []string, peerID
|
||||
}
|
||||
}
|
||||
|
||||
distributionGroupPeers := make([]*ComponentPeer, 0, len(distPeersWithPolicy))
|
||||
distributionGroupPeers := make([]*nbpeer.Peer, 0, len(distPeersWithPolicy))
|
||||
for pID := range distPeersWithPolicy {
|
||||
peer := a.Peers[pID]
|
||||
if peer == nil {
|
||||
continue
|
||||
}
|
||||
distributionGroupPeers = append(distributionGroupPeers, peer.ToComponent())
|
||||
distributionGroupPeers = append(distributionGroupPeers, peer)
|
||||
}
|
||||
return distributionGroupPeers
|
||||
}
|
||||
@@ -1799,66 +1762,3 @@ func filterZoneRecordsForPeers(peer *nbpeer.Peer, customZone nbdns.CustomZone, p
|
||||
|
||||
return filteredRecords
|
||||
}
|
||||
|
||||
// filterPeerAppliedZones filters account zones based on the peer's group membership
|
||||
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
|
||||
}
|
||||
|
||||
@@ -2,7 +2,6 @@ package types
|
||||
|
||||
import (
|
||||
"context"
|
||||
"slices"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
@@ -11,7 +10,6 @@ import (
|
||||
"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
|
||||
@@ -92,6 +90,9 @@ func (a *Account) GetPeerNetworkMapFromComponents(
|
||||
return nm
|
||||
}
|
||||
|
||||
// GetPeerNetworkMapComponents builds the account's slim twin store and computes
|
||||
// the peer's components on it. The calculation itself lives on
|
||||
// networkmap.NetworkMapData and never touches the Account.
|
||||
func (a *Account) GetPeerNetworkMapComponents(
|
||||
ctx context.Context,
|
||||
peerID string,
|
||||
@@ -102,623 +103,10 @@ func (a *Account) GetPeerNetworkMapComponents(
|
||||
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: peer.ToComponent()},
|
||||
})
|
||||
nmd := a.toNetworkMapData(accountZones, validatedPeersMap, resourcePolicies, routers, groupIDToUserIDs)
|
||||
components := nmd.GetPeerNetworkMapComponents(peerID, TwinCustomZone(peersCustomZone))
|
||||
if components != nil {
|
||||
components.ForceRoutingPeerDNSResolution = a.forcesRoutingPeerDNSResolution(peerID, routers)
|
||||
}
|
||||
|
||||
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: peer.ToComponent()},
|
||||
})
|
||||
}
|
||||
|
||||
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)),
|
||||
}
|
||||
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 := a.getPeersGroupsPoliciesRoutes(ctx, peerID, peer.SSHEnabled, validatedPeersMap, &components.PostureFailedPeers)
|
||||
|
||||
if len(sshReqs.neededGroupIDs) > 0 {
|
||||
components.GroupIDToUserIDs = filterGroupIDToUserIDs(groupIDToUserIDs, sshReqs.neededGroupIDs)
|
||||
}
|
||||
if sshReqs.needAllowedUserIDs {
|
||||
components.AllowedUserIDs = a.getAllowedUserIDs()
|
||||
}
|
||||
|
||||
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, 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 = a.getUniquePeerIDsFromGroupsIDs(ctx, policy.SourceGroups())
|
||||
}
|
||||
for _, pID := range a.getPostureValidPeersSaveFailed(peers, policy.SourcePostureChecks, validatedPeersMap, &components.PostureFailedPeers) {
|
||||
if _, exists := components.Peers[pID]; !exists {
|
||||
components.Peers[pID] = a.GetPeer(pID).ToComponent()
|
||||
}
|
||||
}
|
||||
} 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 := a.validatePostureChecksOnPeerGetFailed(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] = g.ToComponent()
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, dstGroupID := range rule.Destinations {
|
||||
if g := a.Groups[dstGroupID]; g != nil {
|
||||
if _, exists := components.Groups[dstGroupID]; !exists {
|
||||
components.Groups[dstGroupID] = g.ToComponent()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
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] = routerTypes.ToComponentMap(networkRoutingPeers)
|
||||
for peerIDKey := range networkRoutingPeers {
|
||||
if p := a.Peers[peerIDKey]; p != nil {
|
||||
cp := components.RouterPeers[peerIDKey]
|
||||
if cp == nil {
|
||||
cp = p.ToComponent()
|
||||
components.RouterPeers[peerIDKey] = cp
|
||||
}
|
||||
if _, exists := components.Peers[peerIDKey]; !exists {
|
||||
if _, validated := validatedPeersMap[peerIDKey]; validated {
|
||||
components.Peers[peerIDKey] = cp
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
components.NetworkResources = append(components.NetworkResources, resource.ToComponent())
|
||||
}
|
||||
}
|
||||
|
||||
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 (a *Account) getPeersGroupsPoliciesRoutes(
|
||||
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] = a.GetPeer(peerID).ToComponent()
|
||||
|
||||
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] = p.ToComponent()
|
||||
}
|
||||
}
|
||||
}
|
||||
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] = p.ToComponent()
|
||||
}
|
||||
}
|
||||
}
|
||||
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 = a.getPeersFromGroups(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 = a.getPeersFromGroups(ctx, rule.Destinations, peerID, nil, validatedPeersMap, postureFailedPeers)
|
||||
}
|
||||
|
||||
if peerInSources {
|
||||
policyRelevant = true
|
||||
for _, pid := range destinationPeers {
|
||||
if _, exists := relevantPeerIDs[pid]; !exists {
|
||||
relevantPeerIDs[pid] = a.GetPeer(pid).ToComponent()
|
||||
}
|
||||
}
|
||||
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] = a.GetPeer(pid).ToComponent()
|
||||
}
|
||||
}
|
||||
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 (a *Account) getPeersFromGroups(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 := a.validatePostureChecksOnPeerGetFailed(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 := a.validatePostureChecksOnPeerGetFailed(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 (a *Account) validatePostureChecksOnPeerGetFailed(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 (a *Account) getPostureValidPeersSaveFailed(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 := a.validatePostureChecksOnPeerGetFailed(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
|
||||
}
|
||||
|
||||
566
management/server/types/account_networkmapdata.go
Normal file
566
management/server/types/account_networkmapdata.go
Normal file
@@ -0,0 +1,566 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"github.com/miekg/dns"
|
||||
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
"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"
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/posture"
|
||||
nbroute "github.com/netbirdio/netbird/route"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
// toNetworkMapData builds the slim twin store from the account once per
|
||||
// account. The per-peer components calculation then runs on the twin.
|
||||
func (a *Account) toNetworkMapData(
|
||||
accountZones []*zones.Zone,
|
||||
validatedPeersMap map[string]struct{},
|
||||
resourcePolicies map[string][]*Policy,
|
||||
routers map[string]map[string]*routerTypes.NetworkRouter,
|
||||
groupIDToUserIDs map[string][]string,
|
||||
) *networkmap.NetworkMapData {
|
||||
nmd := &networkmap.NetworkMapData{
|
||||
Peers: make(map[string]*nmdata.Peer, len(a.Peers)),
|
||||
Groups: make(map[string]*nmdata.Group, len(a.Groups)),
|
||||
Policies: make([]*nmdata.Policy, 0, len(a.Policies)),
|
||||
Routes: make([]*nmdata.Route, 0, len(a.Routes)),
|
||||
NameServerGroups: make([]*nmdata.NameServerGroup, 0, len(a.NameServerGroups)),
|
||||
NetworkResources: make([]*nmdata.NetworkResource, 0, len(a.NetworkResources)),
|
||||
PostureChecks: make(map[string]*nmdata.PostureChecks, len(a.PostureChecks)),
|
||||
ResourcePolicies: make(map[string][]*nmdata.Policy, len(resourcePolicies)),
|
||||
Routers: make(map[string]map[string]*nmdata.NetworkRouter, len(routers)),
|
||||
ValidatedPeers: validatedPeersMap,
|
||||
GroupIDToUserIDs: groupIDToUserIDs,
|
||||
AllowedUserIDs: a.getAllowedUserIDs(),
|
||||
NetworkXIDToPublicID: make(map[string]string, len(a.Networks)),
|
||||
PostureCheckXIDToPublicID: make(map[string]string, len(a.PostureChecks)),
|
||||
}
|
||||
|
||||
if a.Network != nil {
|
||||
nmd.Network = TwinNetwork(a.Network)
|
||||
}
|
||||
nmd.DNSSettings = &nmdata.DNSSettings{DisabledManagementGroups: a.DNSSettings.DisabledManagementGroups}
|
||||
nmd.AccountSettings = TwinAccountSettings(a.Settings)
|
||||
|
||||
for id, p := range a.Peers {
|
||||
nmd.Peers[id] = twinPeer(p)
|
||||
}
|
||||
for id, g := range a.Groups {
|
||||
nmd.Groups[id] = twinGroup(g)
|
||||
}
|
||||
|
||||
policyCache := make(map[string]*nmdata.Policy, len(a.Policies))
|
||||
twinPol := func(p *Policy) *nmdata.Policy {
|
||||
if p == nil {
|
||||
return nil
|
||||
}
|
||||
if tp, ok := policyCache[p.ID]; ok {
|
||||
return tp
|
||||
}
|
||||
tp := twinPolicy(p)
|
||||
policyCache[p.ID] = tp
|
||||
return tp
|
||||
}
|
||||
for _, p := range a.Policies {
|
||||
nmd.Policies = append(nmd.Policies, twinPol(p))
|
||||
}
|
||||
for resID, pols := range resourcePolicies {
|
||||
twinPols := make([]*nmdata.Policy, 0, len(pols))
|
||||
for _, p := range pols {
|
||||
twinPols = append(twinPols, twinPol(p))
|
||||
}
|
||||
nmd.ResourcePolicies[resID] = twinPols
|
||||
}
|
||||
|
||||
for _, r := range a.Routes {
|
||||
if r == nil {
|
||||
continue
|
||||
}
|
||||
nmd.Routes = append(nmd.Routes, twinRoute(r))
|
||||
}
|
||||
for _, nsg := range a.NameServerGroups {
|
||||
nmd.NameServerGroups = append(nmd.NameServerGroups, twinNSG(nsg))
|
||||
}
|
||||
for _, res := range a.NetworkResources {
|
||||
nmd.NetworkResources = append(nmd.NetworkResources, twinNetworkResource(res))
|
||||
}
|
||||
for _, pc := range a.PostureChecks {
|
||||
if pc != nil {
|
||||
nmd.PostureChecks[pc.ID] = twinPostureChecks(pc)
|
||||
nmd.PostureCheckXIDToPublicID[pc.ID] = pc.PublicID
|
||||
}
|
||||
}
|
||||
for _, n := range a.Networks {
|
||||
if n != nil {
|
||||
nmd.NetworkXIDToPublicID[n.ID] = n.PublicID
|
||||
}
|
||||
}
|
||||
for networkID, inner := range routers {
|
||||
twinInner := make(map[string]*nmdata.NetworkRouter, len(inner))
|
||||
for peerID, router := range inner {
|
||||
twinInner[peerID] = twinRouter(router)
|
||||
}
|
||||
nmd.Routers[networkID] = twinInner
|
||||
}
|
||||
|
||||
nmd.AppliedZoneCandidates = buildAppliedZoneCandidates(accountZones)
|
||||
nmd.PrivateServiceCandidates = a.buildPrivateServiceCandidates()
|
||||
|
||||
return nmd
|
||||
}
|
||||
|
||||
func twinPeer(p *nbpeer.Peer) *nmdata.Peer {
|
||||
if p == nil {
|
||||
return nil
|
||||
}
|
||||
networkAddresses := make([]nmdata.NetworkAddress, 0, len(p.Meta.NetworkAddresses))
|
||||
for _, na := range p.Meta.NetworkAddresses {
|
||||
networkAddresses = append(networkAddresses, nmdata.NetworkAddress{NetIP: na.NetIP})
|
||||
}
|
||||
files := make([]nmdata.File, 0, len(p.Meta.Files))
|
||||
for _, f := range p.Meta.Files {
|
||||
files = append(files, nmdata.File{Path: f.Path, ProcessIsRunning: f.ProcessIsRunning})
|
||||
}
|
||||
return &nmdata.Peer{
|
||||
ID: p.ID,
|
||||
Key: p.Key,
|
||||
SSHKey: p.SSHKey,
|
||||
DNSLabel: p.DNSLabel,
|
||||
UserID: p.UserID,
|
||||
SSHEnabled: p.SSHEnabled,
|
||||
LoginExpirationEnabled: p.LoginExpirationEnabled,
|
||||
LastLogin: p.LastLogin,
|
||||
IP: p.IP,
|
||||
IPv6: p.IPv6,
|
||||
RequiresApproval: p.Status != nil && p.Status.RequiresApproval,
|
||||
ExtraDNSLabels: p.ExtraDNSLabels,
|
||||
ProxyMeta: nmdata.ProxyMeta{Embedded: p.ProxyMeta.Embedded},
|
||||
Meta: nmdata.PeerSystemMeta{
|
||||
WtVersion: p.Meta.WtVersion,
|
||||
GoOS: p.Meta.GoOS,
|
||||
OSVersion: p.Meta.OSVersion,
|
||||
KernelVersion: p.Meta.KernelVersion,
|
||||
NetworkAddresses: networkAddresses,
|
||||
Files: files,
|
||||
Capabilities: p.Meta.Capabilities,
|
||||
SyncMessageVersion: p.Meta.SyncMessageVersion,
|
||||
Flags: nmdata.Flags{
|
||||
ServerSSHAllowed: p.Meta.Flags.ServerSSHAllowed,
|
||||
DisableIPv6: p.Meta.Flags.DisableIPv6,
|
||||
},
|
||||
},
|
||||
Location: nmdata.PeerLocation{
|
||||
CountryCode: p.Location.CountryCode,
|
||||
CityName: p.Location.CityName,
|
||||
ConnectionIP: p.Location.ConnectionIP,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// TwinPeer converts a real peer to its slim nmdata twin. Exported for the
|
||||
// port-forwarding integration, which builds proxy NetworkMaps holding twins.
|
||||
func TwinPeer(p *nbpeer.Peer) *nmdata.Peer {
|
||||
return twinPeer(p)
|
||||
}
|
||||
|
||||
// TwinPeers converts real peers to their slim nmdata twins.
|
||||
func TwinPeers(peers []*nbpeer.Peer) []*nmdata.Peer {
|
||||
out := make([]*nmdata.Peer, len(peers))
|
||||
for i, p := range peers {
|
||||
out[i] = twinPeer(p)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// TwinGroups converts real groups to their slim nmdata twins.
|
||||
func TwinGroups(groups []*Group) []*nmdata.Group {
|
||||
out := make([]*nmdata.Group, len(groups))
|
||||
for i, g := range groups {
|
||||
out[i] = twinGroup(g)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func twinGroup(g *Group) *nmdata.Group {
|
||||
if g == nil {
|
||||
return nil
|
||||
}
|
||||
return &nmdata.Group{
|
||||
ID: g.ID,
|
||||
Name: g.Name,
|
||||
PublicID: g.PublicID,
|
||||
Peers: g.Peers,
|
||||
}
|
||||
}
|
||||
|
||||
func twinPolicy(p *Policy) *nmdata.Policy {
|
||||
if p == nil {
|
||||
return nil
|
||||
}
|
||||
rules := make([]*nmdata.PolicyRule, 0, len(p.Rules))
|
||||
for _, r := range p.Rules {
|
||||
rules = append(rules, twinRule(r))
|
||||
}
|
||||
return &nmdata.Policy{
|
||||
ID: p.ID,
|
||||
PublicID: p.PublicID,
|
||||
Enabled: p.Enabled,
|
||||
SourcePostureChecks: p.SourcePostureChecks,
|
||||
Rules: rules,
|
||||
}
|
||||
}
|
||||
|
||||
func twinRule(r *PolicyRule) *nmdata.PolicyRule {
|
||||
if r == nil {
|
||||
return nil
|
||||
}
|
||||
var portRanges []nmdata.RulePortRange
|
||||
if r.PortRanges != nil {
|
||||
portRanges = make([]nmdata.RulePortRange, len(r.PortRanges))
|
||||
for i, pr := range r.PortRanges {
|
||||
portRanges[i] = nmdata.RulePortRange{Start: pr.Start, End: pr.End}
|
||||
}
|
||||
}
|
||||
return &nmdata.PolicyRule{
|
||||
ID: r.ID,
|
||||
PolicyID: r.PolicyID,
|
||||
Enabled: r.Enabled,
|
||||
Action: string(r.Action),
|
||||
Protocol: string(r.Protocol),
|
||||
Bidirectional: r.Bidirectional,
|
||||
Sources: r.Sources,
|
||||
Destinations: r.Destinations,
|
||||
SourceResource: nmdata.Resource{ID: r.SourceResource.ID, Type: string(r.SourceResource.Type)},
|
||||
DestinationResource: nmdata.Resource{ID: r.DestinationResource.ID, Type: string(r.DestinationResource.Type)},
|
||||
Ports: r.Ports,
|
||||
PortRanges: portRanges,
|
||||
AuthorizedGroups: r.AuthorizedGroups,
|
||||
AuthorizedUser: r.AuthorizedUser,
|
||||
}
|
||||
}
|
||||
|
||||
func twinRoute(r *nbroute.Route) *nmdata.Route {
|
||||
return &nmdata.Route{
|
||||
ID: string(r.ID),
|
||||
AccountID: r.AccountID,
|
||||
PublicID: r.PublicID,
|
||||
Network: r.Network,
|
||||
Domains: r.Domains,
|
||||
KeepRoute: r.KeepRoute,
|
||||
NetID: string(r.NetID),
|
||||
Description: r.Description,
|
||||
Peer: r.Peer,
|
||||
PeerID: r.PeerID,
|
||||
PeerGroups: r.PeerGroups,
|
||||
NetworkType: int(r.NetworkType),
|
||||
Masquerade: r.Masquerade,
|
||||
Metric: r.Metric,
|
||||
Enabled: r.Enabled,
|
||||
Groups: r.Groups,
|
||||
AccessControlGroups: r.AccessControlGroups,
|
||||
SkipAutoApply: r.SkipAutoApply,
|
||||
}
|
||||
}
|
||||
|
||||
// TwinRoute converts a real *route.Route to its slim nmdata twin. Exported for
|
||||
// tests that assert against twin routes returned in a NetworkMap.
|
||||
func TwinRoute(r *nbroute.Route) *nmdata.Route {
|
||||
return twinRoute(r)
|
||||
}
|
||||
|
||||
func twinNetworkResource(r *resourceTypes.NetworkResource) *nmdata.NetworkResource {
|
||||
if r == nil {
|
||||
return nil
|
||||
}
|
||||
return &nmdata.NetworkResource{
|
||||
ID: r.ID,
|
||||
NetworkID: r.NetworkID,
|
||||
AccountID: r.AccountID,
|
||||
PublicID: r.PublicID,
|
||||
Name: r.Name,
|
||||
Description: r.Description,
|
||||
Type: string(r.Type),
|
||||
Address: r.Address,
|
||||
Domain: r.Domain,
|
||||
Prefix: r.Prefix,
|
||||
Enabled: r.Enabled,
|
||||
}
|
||||
}
|
||||
|
||||
func twinRouter(r *routerTypes.NetworkRouter) *nmdata.NetworkRouter {
|
||||
if r == nil {
|
||||
return nil
|
||||
}
|
||||
return &nmdata.NetworkRouter{
|
||||
PublicID: r.PublicID,
|
||||
PeerGroups: r.PeerGroups,
|
||||
Masquerade: r.Masquerade,
|
||||
Metric: r.Metric,
|
||||
Enabled: r.Enabled,
|
||||
}
|
||||
}
|
||||
|
||||
func twinNSG(n *nbdns.NameServerGroup) *nmdata.NameServerGroup {
|
||||
if n == nil {
|
||||
return nil
|
||||
}
|
||||
nameServers := make([]nmdata.NameServer, 0, len(n.NameServers))
|
||||
for _, ns := range n.NameServers {
|
||||
nameServers = append(nameServers, nmdata.NameServer{
|
||||
IP: ns.IP,
|
||||
NSType: int(ns.NSType),
|
||||
Port: ns.Port,
|
||||
})
|
||||
}
|
||||
return &nmdata.NameServerGroup{
|
||||
ID: n.ID,
|
||||
PublicID: n.PublicID,
|
||||
Name: n.Name,
|
||||
Description: n.Description,
|
||||
NameServers: nameServers,
|
||||
Groups: n.Groups,
|
||||
Primary: n.Primary,
|
||||
Domains: n.Domains,
|
||||
Enabled: n.Enabled,
|
||||
SearchDomainsEnabled: n.SearchDomainsEnabled,
|
||||
}
|
||||
}
|
||||
|
||||
// TwinNetwork converts a real *Network to its slim twin. Exported for the
|
||||
// graceful-degrade path that builds a minimal NetworkMapComponents directly.
|
||||
func TwinNetwork(n *Network) *nmdata.Network {
|
||||
nc := n.Copy()
|
||||
return &nmdata.Network{
|
||||
Identifier: nc.Identifier,
|
||||
Net: nc.Net,
|
||||
NetV6: nc.NetV6,
|
||||
Dns: nc.Dns,
|
||||
Serial: int64(nc.Serial),
|
||||
}
|
||||
}
|
||||
|
||||
func twinPostureChecks(pc *posture.Checks) *nmdata.PostureChecks {
|
||||
if pc == nil {
|
||||
return nil
|
||||
}
|
||||
out := &nmdata.PostureChecks{ID: pc.ID}
|
||||
def := pc.Checks
|
||||
if def.NBVersionCheck != nil {
|
||||
out.Checks.NBVersionCheck = &nmdata.NBVersionCheck{MinVersion: def.NBVersionCheck.MinVersion}
|
||||
}
|
||||
if def.OSVersionCheck != nil {
|
||||
oc := &nmdata.OSVersionCheck{}
|
||||
if def.OSVersionCheck.Android != nil {
|
||||
oc.Android = &nmdata.MinVersionCheck{MinVersion: def.OSVersionCheck.Android.MinVersion}
|
||||
}
|
||||
if def.OSVersionCheck.Darwin != nil {
|
||||
oc.Darwin = &nmdata.MinVersionCheck{MinVersion: def.OSVersionCheck.Darwin.MinVersion}
|
||||
}
|
||||
if def.OSVersionCheck.Ios != nil {
|
||||
oc.Ios = &nmdata.MinVersionCheck{MinVersion: def.OSVersionCheck.Ios.MinVersion}
|
||||
}
|
||||
if def.OSVersionCheck.Linux != nil {
|
||||
oc.Linux = &nmdata.MinKernelVersionCheck{MinKernelVersion: def.OSVersionCheck.Linux.MinKernelVersion}
|
||||
}
|
||||
if def.OSVersionCheck.Windows != nil {
|
||||
oc.Windows = &nmdata.MinKernelVersionCheck{MinKernelVersion: def.OSVersionCheck.Windows.MinKernelVersion}
|
||||
}
|
||||
out.Checks.OSVersionCheck = oc
|
||||
}
|
||||
if def.GeoLocationCheck != nil {
|
||||
gc := &nmdata.GeoLocationCheck{Action: def.GeoLocationCheck.Action}
|
||||
for _, loc := range def.GeoLocationCheck.Locations {
|
||||
gc.Locations = append(gc.Locations, nmdata.GeoLocation{CountryCode: loc.CountryCode, CityName: loc.CityName})
|
||||
}
|
||||
out.Checks.GeoLocationCheck = gc
|
||||
}
|
||||
if def.PeerNetworkRangeCheck != nil {
|
||||
out.Checks.PeerNetworkRangeCheck = &nmdata.PeerNetworkRangeCheck{
|
||||
Action: def.PeerNetworkRangeCheck.Action,
|
||||
Ranges: def.PeerNetworkRangeCheck.Ranges,
|
||||
}
|
||||
}
|
||||
if def.ProcessCheck != nil {
|
||||
procs := make([]nmdata.Process, 0, len(def.ProcessCheck.Processes))
|
||||
for _, p := range def.ProcessCheck.Processes {
|
||||
procs = append(procs, nmdata.Process{LinuxPath: p.LinuxPath, MacPath: p.MacPath, WindowsPath: p.WindowsPath})
|
||||
}
|
||||
out.Checks.ProcessCheck = &nmdata.ProcessCheck{Processes: procs}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// buildAppliedZoneCandidates precomputes the account-level custom DNS zones
|
||||
// (record conversion) once; the per-peer distribution-group gate runs in the
|
||||
// components calc. Mirrors the account-level half of filterPeerAppliedZones.
|
||||
func buildAppliedZoneCandidates(accountZones []*zones.Zone) []networkmap.AppliedZoneCandidate {
|
||||
var out []networkmap.AppliedZoneCandidate
|
||||
for _, zone := range accountZones {
|
||||
if !zone.Enabled || len(zone.Records) == 0 {
|
||||
continue
|
||||
}
|
||||
simpleRecords := make([]nmdata.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:
|
||||
continue
|
||||
}
|
||||
simpleRecords = append(simpleRecords, nmdata.SimpleRecord{
|
||||
Name: dns.Fqdn(record.Name),
|
||||
Type: recordType,
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: record.TTL,
|
||||
RData: rData,
|
||||
})
|
||||
}
|
||||
out = append(out, networkmap.AppliedZoneCandidate{
|
||||
DistributionGroups: zone.DistributionGroups,
|
||||
Zone: nmdata.CustomZone{
|
||||
Domain: dns.Fqdn(zone.Domain),
|
||||
Records: simpleRecords,
|
||||
SearchDomainDisabled: !zone.EnableSearchDomain,
|
||||
NonAuthoritative: true,
|
||||
},
|
||||
})
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// buildPrivateServiceCandidates precomputes the connected-proxy A records per
|
||||
// private service (account-level); the per-peer access-group gate + apex merge
|
||||
// run in the components calc. Mirrors the account-level half of
|
||||
// SynthesizePrivateServiceZones.
|
||||
func (a *Account) buildPrivateServiceCandidates() []networkmap.PrivateServiceCandidate {
|
||||
if len(a.Services) == 0 {
|
||||
return nil
|
||||
}
|
||||
proxyPeersByCluster := a.GetProxyPeers()
|
||||
if len(proxyPeersByCluster) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
var out []networkmap.PrivateServiceCandidate
|
||||
for _, svc := range a.Services {
|
||||
if svc == nil || !svc.Enabled || !svc.Private {
|
||||
continue
|
||||
}
|
||||
if len(svc.AccessGroups) == 0 {
|
||||
continue
|
||||
}
|
||||
proxyPeers := proxyPeersByCluster[svc.ProxyCluster]
|
||||
if len(proxyPeers) == 0 {
|
||||
continue
|
||||
}
|
||||
apex := a.privateServiceDomainZone(svc)
|
||||
if apex == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
var recs []nmdata.SimpleRecord
|
||||
for _, p := range proxyPeers {
|
||||
if p == nil || !p.IP.IsValid() {
|
||||
continue
|
||||
}
|
||||
if p.Status == nil || !p.Status.Connected {
|
||||
continue
|
||||
}
|
||||
recs = append(recs, nmdata.SimpleRecord{
|
||||
Name: dns.Fqdn(svc.Domain),
|
||||
Type: int(dns.TypeA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: privateServiceDNSRecordTTL,
|
||||
RData: p.IP.String(),
|
||||
})
|
||||
}
|
||||
if len(recs) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
out = append(out, networkmap.PrivateServiceCandidate{
|
||||
AccessGroups: svc.AccessGroups,
|
||||
Zone: nmdata.CustomZone{
|
||||
Domain: dns.Fqdn(apex),
|
||||
Records: recs,
|
||||
NonAuthoritative: true,
|
||||
SearchDomainDisabled: true,
|
||||
},
|
||||
})
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// TwinAccountSettings converts real account settings to the slim nmdata twin.
|
||||
// Exported for callers of the twin-based sync response builders.
|
||||
func TwinAccountSettings(s *Settings) *nmdata.AccountSettingsInfo {
|
||||
if s == nil {
|
||||
return nil
|
||||
}
|
||||
return &nmdata.AccountSettingsInfo{
|
||||
PeerLoginExpirationEnabled: s.PeerLoginExpirationEnabled,
|
||||
PeerLoginExpiration: s.PeerLoginExpiration,
|
||||
PeerInactivityExpirationEnabled: s.PeerInactivityExpirationEnabled,
|
||||
PeerInactivityExpiration: s.PeerInactivityExpiration,
|
||||
DNSDomain: s.DNSDomain,
|
||||
IPv6EnabledGroups: s.IPv6EnabledGroups,
|
||||
RoutingPeerDNSResolutionEnabled: s.RoutingPeerDNSResolutionEnabled,
|
||||
LazyConnectionEnabled: s.LazyConnectionEnabled,
|
||||
AutoUpdateVersion: s.AutoUpdateVersion,
|
||||
AutoUpdateAlways: s.AutoUpdateAlways,
|
||||
MetricsPushEnabled: s.MetricsPushEnabled,
|
||||
}
|
||||
}
|
||||
|
||||
func fromTwinCustomZone(z nmdata.CustomZone) nbdns.CustomZone {
|
||||
records := make([]nbdns.SimpleRecord, 0, len(z.Records))
|
||||
for _, r := range z.Records {
|
||||
records = append(records, nbdns.SimpleRecord{
|
||||
Name: r.Name,
|
||||
Type: r.Type,
|
||||
Class: r.Class,
|
||||
TTL: r.TTL,
|
||||
RData: r.RData,
|
||||
})
|
||||
}
|
||||
return nbdns.CustomZone{
|
||||
Domain: z.Domain,
|
||||
Records: records,
|
||||
SearchDomainDisabled: z.SearchDomainDisabled,
|
||||
NonAuthoritative: z.NonAuthoritative,
|
||||
}
|
||||
}
|
||||
|
||||
// TwinCustomZone converts a real DNS custom zone to its slim nmdata twin.
|
||||
// Exported for the network-map controller's DB-store path, which feeds real
|
||||
// zones into the twin-based components calculation.
|
||||
func TwinCustomZone(z nbdns.CustomZone) nmdata.CustomZone {
|
||||
records := make([]nmdata.SimpleRecord, 0, len(z.Records))
|
||||
for _, r := range z.Records {
|
||||
records = append(records, nmdata.SimpleRecord{
|
||||
Name: r.Name,
|
||||
Type: r.Type,
|
||||
Class: r.Class,
|
||||
TTL: r.TTL,
|
||||
RData: r.RData,
|
||||
})
|
||||
}
|
||||
return nmdata.CustomZone{
|
||||
Domain: z.Domain,
|
||||
Records: records,
|
||||
SearchDomainDisabled: z.SearchDomainDisabled,
|
||||
NonAuthoritative: z.NonAuthoritative,
|
||||
}
|
||||
}
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
func TestPrivateService_NetworkMap_UserPeer_AndProxyPeer(t *testing.T) {
|
||||
@@ -48,7 +49,7 @@ func TestPrivateService_NetworkMap_UserPeer_AndProxyPeer(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func netmapPeerIDs(peers []*ComponentPeer) []string {
|
||||
func netmapPeerIDs(peers []*nmdata.Peer) []string {
|
||||
ids := make([]string, 0, len(peers))
|
||||
for _, p := range peers {
|
||||
if p == nil {
|
||||
|
||||
@@ -13,8 +13,6 @@ import (
|
||||
|
||||
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"
|
||||
networkTypes "github.com/netbirdio/netbird/management/server/networks/types"
|
||||
@@ -666,7 +664,7 @@ func Test_ExpandPortsAndRanges_SSHRuleExpansion(t *testing.T) {
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := ExpandPortsAndRanges(tt.base, tt.rule, tt.peer.ToComponent())
|
||||
result := ExpandPortsAndRanges(tt.base, tt.rule, tt.peer)
|
||||
|
||||
var ports []string
|
||||
for _, fr := range result {
|
||||
@@ -1040,518 +1038,6 @@ func Test_FilterZoneRecordsForPeers(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func Test_filterPeerAppliedZones(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
accountZones []*zones.Zone
|
||||
peerGroups LookupMap
|
||||
expected []nbdns.CustomZone
|
||||
}{
|
||||
{
|
||||
name: "empty peer groups returns empty custom zones",
|
||||
accountZones: []*zones.Zone{},
|
||||
peerGroups: LookupMap{},
|
||||
expected: []nbdns.CustomZone{},
|
||||
},
|
||||
{
|
||||
name: "peer has access to zone with A record",
|
||||
accountZones: []*zones.Zone{
|
||||
{
|
||||
ID: "zone1",
|
||||
Domain: "example.com",
|
||||
Enabled: true,
|
||||
EnableSearchDomain: false,
|
||||
DistributionGroups: []string{"group1"},
|
||||
Records: []*records.Record{
|
||||
{
|
||||
ID: "record1",
|
||||
Name: "www.example.com",
|
||||
Type: records.RecordTypeA,
|
||||
Content: "192.168.1.1",
|
||||
TTL: 300,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
peerGroups: LookupMap{"group1": struct{}{}},
|
||||
expected: []nbdns.CustomZone{
|
||||
{
|
||||
Domain: "example.com.",
|
||||
Records: []nbdns.SimpleRecord{
|
||||
{
|
||||
Name: "www.example.com.",
|
||||
Type: int(dns.TypeA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: 300,
|
||||
RData: "192.168.1.1",
|
||||
},
|
||||
},
|
||||
SearchDomainDisabled: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "peer has access to zone with search domain enabled",
|
||||
accountZones: []*zones.Zone{
|
||||
{
|
||||
ID: "zone1",
|
||||
Domain: "internal.local",
|
||||
Enabled: true,
|
||||
EnableSearchDomain: true,
|
||||
DistributionGroups: []string{"group1"},
|
||||
Records: []*records.Record{
|
||||
{
|
||||
ID: "record1",
|
||||
Name: "api.internal.local",
|
||||
Type: records.RecordTypeA,
|
||||
Content: "10.0.0.1",
|
||||
TTL: 600,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
peerGroups: LookupMap{"group1": struct{}{}},
|
||||
expected: []nbdns.CustomZone{
|
||||
{
|
||||
Domain: "internal.local.",
|
||||
Records: []nbdns.SimpleRecord{
|
||||
{
|
||||
Name: "api.internal.local.",
|
||||
Type: int(dns.TypeA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: 600,
|
||||
RData: "10.0.0.1",
|
||||
},
|
||||
},
|
||||
SearchDomainDisabled: false,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "peer has no access to zone",
|
||||
accountZones: []*zones.Zone{
|
||||
{
|
||||
ID: "zone1",
|
||||
Domain: "private.com",
|
||||
Enabled: true,
|
||||
EnableSearchDomain: false,
|
||||
DistributionGroups: []string{"group2"},
|
||||
Records: []*records.Record{
|
||||
{
|
||||
ID: "record1",
|
||||
Name: "secret.private.com",
|
||||
Type: records.RecordTypeA,
|
||||
Content: "192.168.1.1",
|
||||
TTL: 300,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
peerGroups: LookupMap{"group1": struct{}{}},
|
||||
expected: []nbdns.CustomZone{},
|
||||
},
|
||||
{
|
||||
name: "disabled zone is filtered out",
|
||||
accountZones: []*zones.Zone{
|
||||
{
|
||||
ID: "zone1",
|
||||
Domain: "disabled.com",
|
||||
Enabled: false,
|
||||
EnableSearchDomain: false,
|
||||
DistributionGroups: []string{"group1"},
|
||||
Records: []*records.Record{
|
||||
{
|
||||
ID: "record1",
|
||||
Name: "www.disabled.com",
|
||||
Type: records.RecordTypeA,
|
||||
Content: "192.168.1.1",
|
||||
TTL: 300,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
peerGroups: LookupMap{"group1": struct{}{}},
|
||||
expected: []nbdns.CustomZone{},
|
||||
},
|
||||
{
|
||||
name: "zone with no records is filtered out",
|
||||
accountZones: []*zones.Zone{
|
||||
{
|
||||
ID: "zone1",
|
||||
Domain: "empty.com",
|
||||
Enabled: true,
|
||||
EnableSearchDomain: false,
|
||||
DistributionGroups: []string{"group1"},
|
||||
Records: []*records.Record{},
|
||||
},
|
||||
},
|
||||
peerGroups: LookupMap{"group1": struct{}{}},
|
||||
expected: []nbdns.CustomZone{},
|
||||
},
|
||||
{
|
||||
name: "peer has access via multiple groups",
|
||||
accountZones: []*zones.Zone{
|
||||
{
|
||||
ID: "zone1",
|
||||
Domain: "multi.com",
|
||||
Enabled: true,
|
||||
EnableSearchDomain: false,
|
||||
DistributionGroups: []string{"group1", "group2", "group3"},
|
||||
Records: []*records.Record{
|
||||
{
|
||||
ID: "record1",
|
||||
Name: "www.multi.com",
|
||||
Type: records.RecordTypeA,
|
||||
Content: "192.168.1.1",
|
||||
TTL: 300,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
peerGroups: LookupMap{"group2": struct{}{}},
|
||||
expected: []nbdns.CustomZone{
|
||||
{
|
||||
Domain: "multi.com.",
|
||||
Records: []nbdns.SimpleRecord{
|
||||
{
|
||||
Name: "www.multi.com.",
|
||||
Type: int(dns.TypeA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: 300,
|
||||
RData: "192.168.1.1",
|
||||
},
|
||||
},
|
||||
SearchDomainDisabled: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "multiple zones with mixed access",
|
||||
accountZones: []*zones.Zone{
|
||||
{
|
||||
ID: "zone1",
|
||||
Domain: "allowed.com",
|
||||
Enabled: true,
|
||||
EnableSearchDomain: false,
|
||||
DistributionGroups: []string{"group1"},
|
||||
Records: []*records.Record{
|
||||
{
|
||||
ID: "record1",
|
||||
Name: "www.allowed.com",
|
||||
Type: records.RecordTypeA,
|
||||
Content: "192.168.1.1",
|
||||
TTL: 300,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: "zone2",
|
||||
Domain: "denied.com",
|
||||
Enabled: true,
|
||||
EnableSearchDomain: false,
|
||||
DistributionGroups: []string{"group2"},
|
||||
Records: []*records.Record{
|
||||
{
|
||||
ID: "record2",
|
||||
Name: "www.denied.com",
|
||||
Type: records.RecordTypeA,
|
||||
Content: "192.168.1.2",
|
||||
TTL: 300,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
peerGroups: LookupMap{"group1": struct{}{}},
|
||||
expected: []nbdns.CustomZone{
|
||||
{
|
||||
Domain: "allowed.com.",
|
||||
Records: []nbdns.SimpleRecord{
|
||||
{
|
||||
Name: "www.allowed.com.",
|
||||
Type: int(dns.TypeA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: 300,
|
||||
RData: "192.168.1.1",
|
||||
},
|
||||
},
|
||||
SearchDomainDisabled: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "zone with multiple record types",
|
||||
accountZones: []*zones.Zone{
|
||||
{
|
||||
ID: "zone1",
|
||||
Domain: "mixed.com",
|
||||
Enabled: true,
|
||||
EnableSearchDomain: false,
|
||||
DistributionGroups: []string{"group1"},
|
||||
Records: []*records.Record{
|
||||
{
|
||||
ID: "record1",
|
||||
Name: "www.mixed.com",
|
||||
Type: records.RecordTypeA,
|
||||
Content: "192.168.1.1",
|
||||
TTL: 300,
|
||||
},
|
||||
{
|
||||
ID: "record2",
|
||||
Name: "ipv6.mixed.com",
|
||||
Type: records.RecordTypeAAAA,
|
||||
Content: "2001:db8::1",
|
||||
TTL: 600,
|
||||
},
|
||||
{
|
||||
ID: "record3",
|
||||
Name: "alias.mixed.com",
|
||||
Type: records.RecordTypeCNAME,
|
||||
Content: "www.mixed.com",
|
||||
TTL: 900,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
peerGroups: LookupMap{"group1": struct{}{}},
|
||||
expected: []nbdns.CustomZone{
|
||||
{
|
||||
Domain: "mixed.com.",
|
||||
Records: []nbdns.SimpleRecord{
|
||||
{
|
||||
Name: "www.mixed.com.",
|
||||
Type: int(dns.TypeA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: 300,
|
||||
RData: "192.168.1.1",
|
||||
},
|
||||
{
|
||||
Name: "ipv6.mixed.com.",
|
||||
Type: int(dns.TypeAAAA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: 600,
|
||||
RData: "2001:db8::1",
|
||||
},
|
||||
{
|
||||
Name: "alias.mixed.com.",
|
||||
Type: int(dns.TypeCNAME),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: 900,
|
||||
RData: "www.mixed.com.",
|
||||
},
|
||||
},
|
||||
SearchDomainDisabled: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "multiple zones both accessible",
|
||||
accountZones: []*zones.Zone{
|
||||
{
|
||||
ID: "zone1",
|
||||
Domain: "first.com",
|
||||
Enabled: true,
|
||||
EnableSearchDomain: true,
|
||||
DistributionGroups: []string{"group1"},
|
||||
Records: []*records.Record{
|
||||
{
|
||||
ID: "record1",
|
||||
Name: "www.first.com",
|
||||
Type: records.RecordTypeA,
|
||||
Content: "192.168.1.1",
|
||||
TTL: 300,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: "zone2",
|
||||
Domain: "second.com",
|
||||
Enabled: true,
|
||||
EnableSearchDomain: false,
|
||||
DistributionGroups: []string{"group1"},
|
||||
Records: []*records.Record{
|
||||
{
|
||||
ID: "record2",
|
||||
Name: "www.second.com",
|
||||
Type: records.RecordTypeA,
|
||||
Content: "192.168.1.2",
|
||||
TTL: 600,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
peerGroups: LookupMap{"group1": struct{}{}},
|
||||
expected: []nbdns.CustomZone{
|
||||
{
|
||||
Domain: "first.com.",
|
||||
Records: []nbdns.SimpleRecord{
|
||||
{
|
||||
Name: "www.first.com.",
|
||||
Type: int(dns.TypeA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: 300,
|
||||
RData: "192.168.1.1",
|
||||
},
|
||||
},
|
||||
SearchDomainDisabled: false,
|
||||
},
|
||||
{
|
||||
Domain: "second.com.",
|
||||
Records: []nbdns.SimpleRecord{
|
||||
{
|
||||
Name: "www.second.com.",
|
||||
Type: int(dns.TypeA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: 600,
|
||||
RData: "192.168.1.2",
|
||||
},
|
||||
},
|
||||
SearchDomainDisabled: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "zone with multiple records of same type",
|
||||
accountZones: []*zones.Zone{
|
||||
{
|
||||
ID: "zone1",
|
||||
Domain: "multi-a.com",
|
||||
Enabled: true,
|
||||
EnableSearchDomain: false,
|
||||
DistributionGroups: []string{"group1"},
|
||||
Records: []*records.Record{
|
||||
{
|
||||
ID: "record1",
|
||||
Name: "www.multi-a.com",
|
||||
Type: records.RecordTypeA,
|
||||
Content: "192.168.1.1",
|
||||
TTL: 300,
|
||||
},
|
||||
{
|
||||
ID: "record2",
|
||||
Name: "www.multi-a.com",
|
||||
Type: records.RecordTypeA,
|
||||
Content: "192.168.1.2",
|
||||
TTL: 300,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
peerGroups: LookupMap{"group1": struct{}{}},
|
||||
expected: []nbdns.CustomZone{
|
||||
{
|
||||
Domain: "multi-a.com.",
|
||||
Records: []nbdns.SimpleRecord{
|
||||
{
|
||||
Name: "www.multi-a.com.",
|
||||
Type: int(dns.TypeA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: 300,
|
||||
RData: "192.168.1.1",
|
||||
},
|
||||
{
|
||||
Name: "www.multi-a.com.",
|
||||
Type: int(dns.TypeA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: 300,
|
||||
RData: "192.168.1.2",
|
||||
},
|
||||
},
|
||||
SearchDomainDisabled: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "peer in multiple groups accessing different zones",
|
||||
accountZones: []*zones.Zone{
|
||||
{
|
||||
ID: "zone1",
|
||||
Domain: "zone1.com",
|
||||
Enabled: true,
|
||||
EnableSearchDomain: false,
|
||||
DistributionGroups: []string{"group1"},
|
||||
Records: []*records.Record{
|
||||
{
|
||||
ID: "record1",
|
||||
Name: "www.zone1.com",
|
||||
Type: records.RecordTypeA,
|
||||
Content: "192.168.1.1",
|
||||
TTL: 300,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: "zone2",
|
||||
Domain: "zone2.com",
|
||||
Enabled: true,
|
||||
EnableSearchDomain: false,
|
||||
DistributionGroups: []string{"group2"},
|
||||
Records: []*records.Record{
|
||||
{
|
||||
ID: "record2",
|
||||
Name: "www.zone2.com",
|
||||
Type: records.RecordTypeA,
|
||||
Content: "192.168.1.2",
|
||||
TTL: 300,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
peerGroups: LookupMap{"group1": struct{}{}, "group2": struct{}{}},
|
||||
expected: []nbdns.CustomZone{
|
||||
{
|
||||
Domain: "zone1.com.",
|
||||
Records: []nbdns.SimpleRecord{
|
||||
{
|
||||
Name: "www.zone1.com.",
|
||||
Type: int(dns.TypeA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: 300,
|
||||
RData: "192.168.1.1",
|
||||
},
|
||||
},
|
||||
SearchDomainDisabled: true,
|
||||
},
|
||||
{
|
||||
Domain: "zone2.com.",
|
||||
Records: []nbdns.SimpleRecord{
|
||||
{
|
||||
Name: "www.zone2.com.",
|
||||
Type: int(dns.TypeA),
|
||||
Class: nbdns.DefaultClass,
|
||||
TTL: 300,
|
||||
RData: "192.168.1.2",
|
||||
},
|
||||
},
|
||||
SearchDomainDisabled: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := filterPeerAppliedZones(ctx, tt.accountZones, tt.peerGroups)
|
||||
require.Equal(t, len(tt.expected), len(result), "number of custom zones should match")
|
||||
|
||||
for i, expectedZone := range tt.expected {
|
||||
assert.Equal(t, expectedZone.Domain, result[i].Domain, "domain should match")
|
||||
assert.Equal(t, expectedZone.SearchDomainDisabled, result[i].SearchDomainDisabled, "search domain disabled flag should match")
|
||||
assert.Equal(t, len(expectedZone.Records), len(result[i].Records), "number of records should match")
|
||||
|
||||
for j, expectedRecord := range expectedZone.Records {
|
||||
assert.Equal(t, expectedRecord.Name, result[i].Records[j].Name, "record name should match")
|
||||
assert.Equal(t, expectedRecord.Type, result[i].Records[j].Type, "record type should match")
|
||||
assert.Equal(t, expectedRecord.Class, result[i].Records[j].Class, "record class should match")
|
||||
assert.Equal(t, expectedRecord.TTL, result[i].Records[j].TTL, "record TTL should match")
|
||||
assert.Equal(t, expectedRecord.RData, result[i].Records[j].RData, "record RData should match")
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestInjectPrivateServicePolicies_ProxyPeerGetsInboundRule(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"net"
|
||||
"net/netip"
|
||||
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
nbroute "github.com/netbirdio/netbird/route"
|
||||
sharedtypes "github.com/netbirdio/netbird/shared/management/types"
|
||||
)
|
||||
@@ -17,6 +18,9 @@ type DNSSettings = sharedtypes.DNSSettings
|
||||
|
||||
type FirewallRule = sharedtypes.FirewallRule
|
||||
|
||||
type Group = sharedtypes.Group
|
||||
type GroupPeer = sharedtypes.GroupPeer
|
||||
|
||||
type Network = sharedtypes.Network
|
||||
type NetworkMap = sharedtypes.NetworkMap
|
||||
type ForwardingRule = sharedtypes.ForwardingRule
|
||||
@@ -38,18 +42,6 @@ type RouteFirewallRule = sharedtypes.RouteFirewallRule
|
||||
|
||||
type NetworkMapComponents = sharedtypes.NetworkMapComponents
|
||||
|
||||
type ComponentPeer = sharedtypes.ComponentPeer
|
||||
type ComponentGroup = sharedtypes.ComponentGroup
|
||||
type ComponentRouter = sharedtypes.ComponentRouter
|
||||
type ComponentResource = sharedtypes.ComponentResource
|
||||
type ComponentResourceType = sharedtypes.ComponentResourceType
|
||||
|
||||
const (
|
||||
ComponentResourceHost = sharedtypes.ComponentResourceHost
|
||||
ComponentResourceSubnet = sharedtypes.ComponentResourceSubnet
|
||||
ComponentResourceDomain = sharedtypes.ComponentResourceDomain
|
||||
)
|
||||
|
||||
var EmptyNetworkMapComponents = sharedtypes.EmptyNetworkMapComponents
|
||||
|
||||
type AccountSettingsInfo = sharedtypes.AccountSettingsInfo
|
||||
@@ -60,7 +52,12 @@ type NetworkMapComponentsCompact = sharedtypes.NetworkMapComponentsCompact
|
||||
type LookupMap = sharedtypes.LookupMap
|
||||
type FirewallRuleContext = sharedtypes.FirewallRuleContext
|
||||
|
||||
const GroupAllName = sharedtypes.GroupAllName
|
||||
const (
|
||||
GroupIssuedAPI = sharedtypes.GroupIssuedAPI
|
||||
GroupIssuedJWT = sharedtypes.GroupIssuedJWT
|
||||
GroupIssuedIntegration = sharedtypes.GroupIssuedIntegration
|
||||
GroupAllName = sharedtypes.GroupAllName
|
||||
)
|
||||
|
||||
// Function forwarders preserve types.X(...) call sites that previously
|
||||
// resolved to package-local funcs. Plain forwarders (not var aliases) keep
|
||||
@@ -70,20 +67,24 @@ func PolicyRuleImpliesLegacySSH(rule *PolicyRule) bool {
|
||||
return sharedtypes.PolicyRuleImpliesLegacySSH(rule)
|
||||
}
|
||||
|
||||
func ExpandPortsAndRanges(base FirewallRule, rule *PolicyRule, peer *ComponentPeer) []*FirewallRule {
|
||||
return sharedtypes.ExpandPortsAndRanges(base, rule, peer)
|
||||
// ExpandPortsAndRanges / AppendIPv6FirewallRule / GenerateRouteFirewallRules
|
||||
// forward to the shared twin-typed helpers, converting the real types the
|
||||
// legacy Account calc still uses to nmdata twins at this boundary.
|
||||
|
||||
func ExpandPortsAndRanges(base FirewallRule, rule *PolicyRule, peer *nbpeer.Peer) []*FirewallRule {
|
||||
return sharedtypes.ExpandPortsAndRanges(base, twinRule(rule), twinPeer(peer))
|
||||
}
|
||||
|
||||
func AppendIPv6FirewallRule(rules []*FirewallRule, rulesExists map[string]struct{}, peer, targetPeer *ComponentPeer, rule *PolicyRule, rc FirewallRuleContext) []*FirewallRule {
|
||||
return sharedtypes.AppendIPv6FirewallRule(rules, rulesExists, peer, targetPeer, rule, rc)
|
||||
func AppendIPv6FirewallRule(rules []*FirewallRule, rulesExists map[string]struct{}, peer, targetPeer *nbpeer.Peer, rule *PolicyRule, rc FirewallRuleContext) []*FirewallRule {
|
||||
return sharedtypes.AppendIPv6FirewallRule(rules, rulesExists, twinPeer(peer), twinPeer(targetPeer), twinRule(rule), rc)
|
||||
}
|
||||
|
||||
func CalculateNetworkMapFromComponents(ctx context.Context, components *NetworkMapComponents) *NetworkMap {
|
||||
return sharedtypes.CalculateNetworkMapFromComponents(ctx, components)
|
||||
}
|
||||
|
||||
func GenerateRouteFirewallRules(ctx context.Context, route *nbroute.Route, rule *PolicyRule, groupPeers []*ComponentPeer, direction int, includeIPv6 bool) []*RouteFirewallRule {
|
||||
return sharedtypes.GenerateRouteFirewallRules(ctx, route, rule, groupPeers, direction, includeIPv6)
|
||||
func GenerateRouteFirewallRules(ctx context.Context, route *nbroute.Route, rule *PolicyRule, groupPeers []*nbpeer.Peer, direction int, includeIPv6 bool) []*RouteFirewallRule {
|
||||
return sharedtypes.GenerateRouteFirewallRules(ctx, twinRoute(route), twinRule(rule), TwinPeers(groupPeers), direction, includeIPv6)
|
||||
}
|
||||
|
||||
func AllocateIPv6Subnet(r *rand.Rand) net.IPNet {
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
func TestNetworkMapComponents_IPv6EndToEnd(t *testing.T) {
|
||||
@@ -105,7 +105,7 @@ func TestNetworkMapComponents_RemotePeerWithoutCapability(t *testing.T) {
|
||||
require.NotNil(t, nm)
|
||||
|
||||
t.Run("AllowedIPs include remote v6", func(t *testing.T) {
|
||||
var dst *types.ComponentPeer
|
||||
var dst *nmdata.Peer
|
||||
for _, p := range nm.Peers {
|
||||
if p.ID == "peer-dst-1" {
|
||||
dst = p
|
||||
|
||||
703
management/server/types/legacynmap/account_components.go
Normal file
703
management/server/types/legacynmap/account_components.go
Normal file
@@ -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
|
||||
}
|
||||
48
management/server/types/legacynmap/aliases.go
Normal file
48
management/server/types/legacynmap/aliases.go
Normal file
@@ -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
|
||||
)
|
||||
@@ -1,4 +1,6 @@
|
||||
package types
|
||||
//go:build nmapequiv
|
||||
|
||||
package legacynmap
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
128
management/server/types/legacynmap/converters.go
Normal file
128
management/server/types/legacynmap/converters.go
Normal file
@@ -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,
|
||||
}
|
||||
}
|
||||
284
management/server/types/legacynmap/copied_funcs.go
Normal file
284
management/server/types/legacynmap/copied_funcs.go
Normal file
@@ -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
|
||||
}
|
||||
7
management/server/types/legacynmap/doc.go
Normal file
7
management/server/types/legacynmap/doc.go
Normal file
@@ -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
|
||||
570
management/server/types/legacynmap/equivalence_test.go
Normal file
570
management/server/types/legacynmap/equivalence_test.go
Normal file
@@ -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 "<nil>"
|
||||
}
|
||||
s := prototext.Format(m)
|
||||
const maxLen = 800
|
||||
if len(s) > maxLen {
|
||||
return s[:maxLen] + "...(truncated)"
|
||||
}
|
||||
return s
|
||||
}
|
||||
157
management/server/types/legacynmap/firewall_helpers.go
Normal file
157
management/server/types/legacynmap/firewall_helpers.go
Normal file
@@ -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]
|
||||
}
|
||||
1034
management/server/types/legacynmap/networkmap_components.go
Normal file
1034
management/server/types/legacynmap/networkmap_components.go
Normal file
File diff suppressed because it is too large
Load Diff
208
management/server/types/legacynmap/proto_legacy.go
Normal file
208
management/server/types/legacynmap/proto_legacy.go
Normal file
@@ -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)
|
||||
}
|
||||
@@ -18,6 +18,7 @@ import (
|
||||
"github.com/netbirdio/netbird/management/server/posture"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/route"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
func networkMapFromComponents(t *testing.T, account *types.Account, peerID string, validatedPeers map[string]struct{}) *types.NetworkMap {
|
||||
@@ -49,7 +50,7 @@ func allPeersValidated(account *types.Account, excludePeerIDs ...string) map[str
|
||||
return validated
|
||||
}
|
||||
|
||||
func peerIDs(peers []*types.ComponentPeer) []string {
|
||||
func peerIDs(peers []*nmdata.Peer) []string {
|
||||
ids := make([]string, len(peers))
|
||||
for i, p := range peers {
|
||||
ids[i] = p.ID
|
||||
@@ -625,7 +626,7 @@ func TestNetworkMapComponents_DomainNetworkResource(t *testing.T) {
|
||||
|
||||
var hasDomainRoute bool
|
||||
for _, r := range nm.Routes {
|
||||
if r.NetworkType == route.DomainNetwork && len(r.Domains) > 0 && r.Domains[0].SafeString() == "api.example.com" {
|
||||
if r.NetworkType == int(route.DomainNetwork) && len(r.Domains) > 0 && r.Domains[0].SafeString() == "api.example.com" {
|
||||
hasDomainRoute = true
|
||||
}
|
||||
}
|
||||
|
||||
@@ -66,7 +66,7 @@ func BenchmarkNetworkMapWireEncode(b *testing.B) {
|
||||
|
||||
// Pre-encode once so the size metric is identical for every run inside
|
||||
// the same scale; the b.Loop call only re-runs encode + Marshal.
|
||||
legacyResp := mgmtgrpc.ToSyncResponse(ctx, nil, nil, nil, peer, nil, nil, networkMap, "netbird.cloud", nil, dnsCache, settings, nil, nil, 0)
|
||||
legacyResp := mgmtgrpc.ToSyncResponse(ctx, nil, nil, nil, types.TwinPeer(peer), nil, nil, networkMap, "netbird.cloud", nil, dnsCache, types.TwinAccountSettings(settings), nil, nil, 0)
|
||||
legacyBytes, err := goproto.Marshal(legacyResp.NetworkMap)
|
||||
if err != nil {
|
||||
b.Fatalf("marshal legacy networkmap: %v", err)
|
||||
@@ -88,7 +88,7 @@ func BenchmarkNetworkMapWireEncode(b *testing.B) {
|
||||
b.ReportMetric(float64(len(legacyBytes)), "bytes/msg")
|
||||
b.ResetTimer()
|
||||
for range b.N {
|
||||
resp := mgmtgrpc.ToSyncResponse(ctx, nil, nil, nil, peer, nil, nil, networkMap, "netbird.cloud", nil, dnsCache, settings, nil, nil, 0)
|
||||
resp := mgmtgrpc.ToSyncResponse(ctx, nil, nil, nil, types.TwinPeer(peer), nil, nil, networkMap, "netbird.cloud", nil, dnsCache, types.TwinAccountSettings(settings), nil, nil, 0)
|
||||
if _, err := goproto.Marshal(resp.NetworkMap); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
@@ -135,7 +135,7 @@ func BenchmarkNetworkMapWireSize(b *testing.B) {
|
||||
dnsCache := &cache.DNSConfigCache{}
|
||||
settings := &types.Settings{}
|
||||
|
||||
legacyResp := mgmtgrpc.ToSyncResponse(ctx, nil, nil, nil, peer, nil, nil, networkMap, "netbird.cloud", nil, dnsCache, settings, nil, nil, 0)
|
||||
legacyResp := mgmtgrpc.ToSyncResponse(ctx, nil, nil, nil, types.TwinPeer(peer), nil, nil, networkMap, "netbird.cloud", nil, dnsCache, types.TwinAccountSettings(settings), nil, nil, 0)
|
||||
legacyBytes, err := goproto.Marshal(legacyResp.NetworkMap)
|
||||
if err != nil {
|
||||
b.Fatalf("marshal legacy networkmap: %v", err)
|
||||
|
||||
@@ -45,7 +45,7 @@ func TestNetworkMapWireBreakdown(t *testing.T) {
|
||||
dnsCache := &cache.DNSConfigCache{}
|
||||
settings := &types.Settings{}
|
||||
|
||||
legacyResp := mgmtgrpc.ToSyncResponse(ctx, nil, nil, nil, peer, nil, nil, networkMap, "netbird.cloud", nil, dnsCache, settings, nil, nil, 0)
|
||||
legacyResp := mgmtgrpc.ToSyncResponse(ctx, nil, nil, nil, types.TwinPeer(peer), nil, nil, networkMap, "netbird.cloud", nil, dnsCache, types.TwinAccountSettings(settings), nil, nil, 0)
|
||||
legacyTotal := mustMarshalSize(t, legacyResp.NetworkMap)
|
||||
|
||||
envelope := mgmtgrpc.EncodeNetworkMapEnvelope(mgmtgrpc.ComponentsEnvelopeInput{
|
||||
|
||||
@@ -19,3 +19,34 @@ func Difference(a, b []string) []string {
|
||||
func ToPtr[T any](value T) *T {
|
||||
return &value
|
||||
}
|
||||
|
||||
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 !contains(result, item) {
|
||||
result = append(result, item)
|
||||
}
|
||||
}
|
||||
|
||||
for _, item := range arr2 {
|
||||
if !contains(result, item) {
|
||||
result = append(result, item)
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
func contains[T comparableObject[T]](slice []T, element T) bool {
|
||||
for _, item := range slice {
|
||||
if item.Equal(element) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
package types
|
||||
package util
|
||||
|
||||
import (
|
||||
"testing"
|
||||
@@ -17,7 +17,7 @@ func (t testObject) Equal(other testObject) bool {
|
||||
func Test_MergeUniqueArraysWithoutDuplicates(t *testing.T) {
|
||||
arr1 := []testObject{{value: 1}, {value: 2}}
|
||||
arr2 := []testObject{{value: 2}, {value: 3}}
|
||||
result := mergeUnique(arr1, arr2)
|
||||
result := MergeUnique(arr1, arr2)
|
||||
assert.Len(t, result, 3)
|
||||
assert.Contains(t, result, testObject{value: 1})
|
||||
assert.Contains(t, result, testObject{value: 2})
|
||||
@@ -27,14 +27,14 @@ func Test_MergeUniqueArraysWithoutDuplicates(t *testing.T) {
|
||||
func Test_MergeUniqueHandlesEmptyArrays(t *testing.T) {
|
||||
arr1 := []testObject{}
|
||||
arr2 := []testObject{}
|
||||
result := mergeUnique(arr1, arr2)
|
||||
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)
|
||||
result := MergeUnique(arr1, arr2)
|
||||
assert.Len(t, result, 2)
|
||||
assert.Contains(t, result, testObject{value: 1})
|
||||
assert.Contains(t, result, testObject{value: 2})
|
||||
@@ -126,7 +126,7 @@ func startManagement(t *testing.T) (*grpc.Server, net.Listener) {
|
||||
|
||||
updateManager := update_channel.NewPeersUpdateManager(metrics)
|
||||
requestBuffer := mgmt.NewAccountRequestBuffer(ctx, store)
|
||||
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peersManger), config)
|
||||
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peersManger), config, nil)
|
||||
accountManager, err := mgmt.BuildManager(context.Background(), config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, ia, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManagerMock, false, cacheStore)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
|
||||
@@ -5,14 +5,15 @@ import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
nbroute "github.com/netbirdio/netbird/route"
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/shared/management/domain"
|
||||
nmdata "github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
"github.com/netbirdio/netbird/shared/management/types"
|
||||
)
|
||||
@@ -35,28 +36,28 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents,
|
||||
Network: decodeAccountNetwork(full.Network),
|
||||
AccountSettings: decodeAccountSettings(full.AccountSettings),
|
||||
CustomZoneDomain: full.CustomZoneDomain,
|
||||
Peers: make(map[string]*types.ComponentPeer, len(full.Peers)),
|
||||
Groups: make(map[string]*types.ComponentGroup, len(full.Groups)),
|
||||
Policies: make([]*types.Policy, 0, len(full.Policies)),
|
||||
Routes: make([]*nbroute.Route, 0, len(full.Routes)),
|
||||
NameServerGroups: make([]*nbdns.NameServerGroup, 0, len(full.NameserverGroups)),
|
||||
Peers: make(map[string]*nmdata.Peer, len(full.Peers)),
|
||||
Groups: make(map[string]*nmdata.Group, len(full.Groups)),
|
||||
Policies: make([]*nmdata.Policy, 0, len(full.Policies)),
|
||||
Routes: make([]*nmdata.Route, 0, len(full.Routes)),
|
||||
NameServerGroups: make([]*nmdata.NameServerGroup, 0, len(full.NameserverGroups)),
|
||||
AllDNSRecords: decodeSimpleRecords(full.AllDnsRecords),
|
||||
AccountZones: decodeCustomZones(full.AccountZones),
|
||||
ResourcePoliciesMap: make(map[string][]*types.Policy),
|
||||
RoutersMap: make(map[string]map[string]*types.ComponentRouter),
|
||||
NetworkResources: make([]*types.ComponentResource, 0, len(full.NetworkResources)),
|
||||
RouterPeers: make(map[string]*types.ComponentPeer),
|
||||
ResourcePoliciesMap: make(map[string][]*nmdata.Policy),
|
||||
RoutersMap: make(map[string]map[string]*nmdata.NetworkRouter),
|
||||
NetworkResources: make([]*nmdata.NetworkResource, 0, len(full.NetworkResources)),
|
||||
RouterPeers: make(map[string]*nmdata.Peer),
|
||||
AllowedUserIDs: stringSliceToSet(full.AllowedUserIds),
|
||||
PostureFailedPeers: make(map[string]map[string]struct{}, len(full.PostureFailedPeers)),
|
||||
GroupIDToUserIDs: make(map[string][]string, len(full.GroupIdToUserIds)),
|
||||
}
|
||||
|
||||
if full.DnsSettings != nil {
|
||||
c.DNSSettings = &types.DNSSettings{
|
||||
c.DNSSettings = &nmdata.DNSSettings{
|
||||
DisabledManagementGroups: full.DnsSettings.DisabledManagementGroupIds,
|
||||
}
|
||||
} else {
|
||||
c.DNSSettings = &types.DNSSettings{}
|
||||
c.DNSSettings = &nmdata.DNSSettings{}
|
||||
}
|
||||
|
||||
// Phase 1: peers. The envelope's peers slice is index-addressed on the
|
||||
@@ -98,10 +99,21 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents,
|
||||
log.WithField("peer idx", idx).Error("unrecognized peer idx during decoding")
|
||||
}
|
||||
}
|
||||
group := &types.ComponentGroup{
|
||||
ID: groupID,
|
||||
PublicID: gc.Id,
|
||||
Peers: peerIDs,
|
||||
|
||||
fromCompactResources := func() []nmdata.Resource {
|
||||
var toret []nmdata.Resource
|
||||
|
||||
for _, r := range gc.Resources {
|
||||
toret = append(toret, resourceFromProto(r, peerIDByIndex))
|
||||
}
|
||||
|
||||
return toret
|
||||
}
|
||||
|
||||
group := &nmdata.Group{
|
||||
PublicID: gc.Id,
|
||||
Peers: peerIDs,
|
||||
Resources: fromCompactResources(),
|
||||
}
|
||||
if gc.IsAll {
|
||||
group.Name = types.GroupAllName
|
||||
@@ -111,7 +123,7 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents,
|
||||
|
||||
// Phase 3: policies (PolicyCompact = one rule per entry; current data
|
||||
// model is 1 rule per policy).
|
||||
policyByID := make(map[string]*types.Policy, len(full.Policies))
|
||||
policyByID := make(map[string]*nmdata.Policy, len(full.Policies))
|
||||
for i, pc := range full.Policies {
|
||||
if pc == nil {
|
||||
return nil, fmt.Errorf("invalid envelope: policies[%d] is nil", i)
|
||||
@@ -148,7 +160,7 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents,
|
||||
// Phase 7: routers_map (outer key = network seq id, inner key = peer-id
|
||||
// reconstructed from peer_index). Synthesized network id is "net_<seq>".
|
||||
for networkID, list := range full.RoutersMap {
|
||||
inner := make(map[string]*types.ComponentRouter, len(list.Entries))
|
||||
inner := make(map[string]*nmdata.NetworkRouter, len(list.Entries))
|
||||
for _, entry := range list.Entries {
|
||||
if !entry.PeerIndexSet {
|
||||
continue
|
||||
@@ -158,10 +170,8 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents,
|
||||
continue
|
||||
}
|
||||
peerID := peerIDByIndex[entry.PeerIndex]
|
||||
inner[peerID] = &types.ComponentRouter{
|
||||
NetworkID: networkID,
|
||||
inner[peerID] = &nmdata.NetworkRouter{
|
||||
PublicID: entry.Id,
|
||||
Peer: peerID,
|
||||
PeerGroups: entry.PeerGroupIds,
|
||||
Masquerade: entry.Masquerade,
|
||||
Metric: int(entry.Metric),
|
||||
@@ -180,7 +190,7 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents,
|
||||
if len(ids.Ids) == 0 {
|
||||
continue
|
||||
}
|
||||
policies := make([]*types.Policy, 0, len(ids.Ids))
|
||||
policies := make([]*nmdata.Policy, 0, len(ids.Ids))
|
||||
for _, id := range ids.Ids {
|
||||
if p, ok := policyByID[id]; ok {
|
||||
policies = append(policies, p)
|
||||
@@ -193,6 +203,15 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents,
|
||||
}
|
||||
}
|
||||
|
||||
// Phase 8: rebuild resource_policies_map
|
||||
for _, r := range c.NetworkResources {
|
||||
policies := policiesForNetworkResource(r.ID, c.Policies, c.Groups)
|
||||
if len(policies) == 0 {
|
||||
continue
|
||||
}
|
||||
c.ResourcePoliciesMap[r.ID] = policies
|
||||
}
|
||||
|
||||
// Phase 9: group_id_to_user_ids — wire keys are seq ids, synth to strings.
|
||||
for groupId, list := range full.GroupIdToUserIds {
|
||||
c.GroupIDToUserIDs[groupId] = append([]string(nil), list.UserIds...)
|
||||
@@ -228,14 +247,51 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents,
|
||||
return c, nil
|
||||
}
|
||||
|
||||
func decodeAccountNetwork(an *proto.AccountNetwork) *types.Network {
|
||||
func networkResourceGroups(resourceId string, groups map[string]*nmdata.Group) []string {
|
||||
var toret []string
|
||||
for _, group := range groups {
|
||||
for _, resource := range group.Resources {
|
||||
if resource.ID == resourceId {
|
||||
toret = append(toret, group.PublicID)
|
||||
}
|
||||
}
|
||||
}
|
||||
return toret
|
||||
}
|
||||
|
||||
func policiesForNetworkResource(resourceId string, allPolicies []*nmdata.Policy, groups map[string]*nmdata.Group) []*nmdata.Policy {
|
||||
var toret []*nmdata.Policy
|
||||
|
||||
networkResourceGroups := networkResourceGroups(resourceId, groups)
|
||||
for _, p := range allPolicies {
|
||||
if p == nil || !p.Enabled {
|
||||
continue
|
||||
}
|
||||
|
||||
// there's always only one rule in each policy
|
||||
if p.Rules[0].DestinationResource.ID == resourceId {
|
||||
toret = append(toret, p)
|
||||
continue
|
||||
}
|
||||
for _, groupId := range networkResourceGroups {
|
||||
if slices.Contains(p.Rules[0].Destinations, groupId) {
|
||||
toret = append(toret, p)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return toret
|
||||
}
|
||||
|
||||
func decodeAccountNetwork(an *proto.AccountNetwork) *nmdata.Network {
|
||||
if an == nil {
|
||||
return nil
|
||||
}
|
||||
n := &types.Network{
|
||||
n := &nmdata.Network{
|
||||
Identifier: an.Identifier,
|
||||
Dns: an.Dns,
|
||||
Serial: an.Serial,
|
||||
Serial: int64(an.Serial),
|
||||
}
|
||||
if an.NetCidr != "" {
|
||||
if _, ipnet, err := net.ParseCIDR(an.NetCidr); err == nil && ipnet != nil {
|
||||
@@ -250,32 +306,50 @@ func decodeAccountNetwork(an *proto.AccountNetwork) *types.Network {
|
||||
return n
|
||||
}
|
||||
|
||||
func decodeAccountSettings(as *proto.AccountSettingsCompact) *types.AccountSettingsInfo {
|
||||
func decodeAccountSettings(as *proto.AccountSettingsCompact) *nmdata.AccountSettingsInfo {
|
||||
if as == nil {
|
||||
return &types.AccountSettingsInfo{}
|
||||
return &nmdata.AccountSettingsInfo{}
|
||||
}
|
||||
return &types.AccountSettingsInfo{
|
||||
return &nmdata.AccountSettingsInfo{
|
||||
PeerLoginExpirationEnabled: as.PeerLoginExpirationEnabled,
|
||||
PeerLoginExpiration: time.Duration(as.PeerLoginExpirationNs),
|
||||
}
|
||||
}
|
||||
|
||||
func decodePeerCompact(pc *proto.PeerCompact, peerID string) *types.ComponentPeer {
|
||||
peer := &types.ComponentPeer{
|
||||
func decodePeerCompact(pc *proto.PeerCompact, peerID string) *nmdata.Peer {
|
||||
var caps []int32
|
||||
if pc.SupportsSourcePrefixes {
|
||||
caps = append(caps, nbpeer.PeerCapabilitySourcePrefixes)
|
||||
}
|
||||
if pc.SupportsIpv6 {
|
||||
caps = append(caps, nbpeer.PeerCapabilityIPv6Overlay)
|
||||
}
|
||||
peer := &nmdata.Peer{
|
||||
ID: peerID,
|
||||
Key: peerID,
|
||||
SSHKey: string(pc.SshPubKey),
|
||||
SSHEnabled: pc.SshEnabled,
|
||||
DNSLabel: pc.DnsLabel,
|
||||
LoginExpirationEnabled: pc.LoginExpirationEnabled,
|
||||
AgentVersion: pc.AgentVersion,
|
||||
SupportsSourcePrefixes: pc.SupportsSourcePrefixes,
|
||||
SupportsIPv6: pc.SupportsIpv6,
|
||||
ServerSSHAllowed: pc.ServerSshAllowed,
|
||||
AddedWithSSOLogin: pc.AddedWithSsoLogin,
|
||||
Meta: nmdata.PeerSystemMeta{
|
||||
WtVersion: pc.AgentVersion,
|
||||
Capabilities: caps,
|
||||
Flags: nmdata.Flags{
|
||||
ServerSSHAllowed: pc.ServerSshAllowed,
|
||||
},
|
||||
},
|
||||
}
|
||||
if pc.AddedWithSsoLogin {
|
||||
// Set a non-empty UserID so (*Peer).AddedWithSSOLogin() returns true.
|
||||
// The original UserID isn't on the wire; the value is intentionally
|
||||
// visibly synthetic so any future consumer that mistakes UserID for a
|
||||
// real account user xid won't silently match (or worse, write the
|
||||
// sentinel into a downstream record).
|
||||
peer.UserID = "<env-sso>"
|
||||
}
|
||||
if pc.LastLoginUnixNano != 0 {
|
||||
peer.LastLogin = time.Unix(0, pc.LastLoginUnixNano)
|
||||
t := time.Unix(0, pc.LastLoginUnixNano)
|
||||
peer.LastLogin = &t
|
||||
}
|
||||
switch len(pc.Ip) {
|
||||
case 4:
|
||||
@@ -293,13 +367,13 @@ func decodePeerCompact(pc *proto.PeerCompact, peerID string) *types.ComponentPee
|
||||
return peer
|
||||
}
|
||||
|
||||
func decodePolicyCompact(pc *proto.PolicyCompact, policyID string, peerIDByIndex []string) *types.Policy {
|
||||
rule := &types.PolicyRule{
|
||||
func decodePolicyCompact(pc *proto.PolicyCompact, policyID string, peerIDByIndex []string) *nmdata.Policy {
|
||||
rule := &nmdata.PolicyRule{
|
||||
ID: policyID, // 1 rule per policy → reuse synthesized id
|
||||
PolicyID: policyID,
|
||||
Enabled: true,
|
||||
Action: actionFromProto(pc.Action),
|
||||
Protocol: protocolFromProto(pc.Protocol),
|
||||
Action: string(actionFromProto(pc.Action)),
|
||||
Protocol: string(protocolFromProto(pc.Protocol)),
|
||||
Bidirectional: pc.Bidirectional,
|
||||
Ports: uint32SliceToStrings(pc.Ports),
|
||||
PortRanges: portRangesFromProto(pc.PortRanges),
|
||||
@@ -310,11 +384,11 @@ func decodePolicyCompact(pc *proto.PolicyCompact, policyID string, peerIDByIndex
|
||||
SourceResource: resourceFromProto(pc.SourceResource, peerIDByIndex),
|
||||
DestinationResource: resourceFromProto(pc.DestinationResource, peerIDByIndex),
|
||||
}
|
||||
return &types.Policy{
|
||||
return &nmdata.Policy{
|
||||
ID: policyID,
|
||||
PublicID: pc.Id,
|
||||
Enabled: true,
|
||||
Rules: []*types.PolicyRule{rule},
|
||||
Rules: []*nmdata.PolicyRule{rule},
|
||||
SourcePostureChecks: pc.SourcePostureCheckIds,
|
||||
}
|
||||
}
|
||||
@@ -322,15 +396,31 @@ func decodePolicyCompact(pc *proto.PolicyCompact, policyID string, peerIDByIndex
|
||||
// resourceFromProto rebuilds types.Resource. For peer-typed resources the
|
||||
// peer reference is reconstructed from the envelope's peer index — wire
|
||||
// format ships no xid for peers, so we use the synthesized peer id.
|
||||
func resourceFromProto(r *proto.ResourceCompact, peerIDByIndex []string) types.Resource {
|
||||
func resourceFromProto(r *proto.ResourceCompact, peerIDByIndex []string) nmdata.Resource {
|
||||
if r == nil {
|
||||
return types.Resource{}
|
||||
return nmdata.Resource{}
|
||||
}
|
||||
out := types.Resource{Type: types.ResourceType(r.Type)}
|
||||
if r.PeerIndexSet && int(r.PeerIndex) < len(peerIDByIndex) {
|
||||
out.ID = peerIDByIndex[r.PeerIndex]
|
||||
|
||||
t, ok := proto.ResourceCompactType_name[int32(r.Type)]
|
||||
if !ok || r.Type == proto.ResourceCompactType_unknown_type {
|
||||
return nmdata.Resource{}
|
||||
}
|
||||
|
||||
if r.Type == proto.ResourceCompactType_peer && int(r.GetPeerIndex()) >= len(peerIDByIndex) {
|
||||
return nmdata.Resource{}
|
||||
}
|
||||
|
||||
if r.Type == proto.ResourceCompactType_peer && int(r.GetPeerIndex()) < len(peerIDByIndex) {
|
||||
return nmdata.Resource{
|
||||
Type: "peer",
|
||||
ID: peerIDByIndex[int(r.GetPeerIndex())],
|
||||
}
|
||||
}
|
||||
|
||||
return nmdata.Resource{
|
||||
Type: t,
|
||||
ID: r.GetId(),
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// authorizedGroupsFromProto inverts encodeAuthorizedGroups: the wire form
|
||||
@@ -351,15 +441,15 @@ func authorizedGroupsFromProto(m map[string]*proto.UserNameList) map[string][]st
|
||||
return out
|
||||
}
|
||||
|
||||
func decodeRouteRaw(rr *proto.RouteRaw, peerIDByIndex []string) *nbroute.Route {
|
||||
r := &nbroute.Route{
|
||||
ID: nbroute.ID(rr.Id),
|
||||
func decodeRouteRaw(rr *proto.RouteRaw, peerIDByIndex []string) *nmdata.Route {
|
||||
r := &nmdata.Route{
|
||||
ID: rr.Id,
|
||||
PublicID: rr.Id,
|
||||
NetID: nbroute.NetID(rr.NetId),
|
||||
NetID: rr.NetId,
|
||||
Description: rr.Description,
|
||||
Domains: domainsFromPunycode(rr.Domains),
|
||||
KeepRoute: rr.KeepRoute,
|
||||
NetworkType: nbroute.NetworkType(rr.NetworkType),
|
||||
NetworkType: int(rr.NetworkType),
|
||||
Masquerade: rr.Masquerade,
|
||||
Metric: int(rr.Metric),
|
||||
Enabled: rr.Enabled,
|
||||
@@ -379,8 +469,8 @@ func decodeRouteRaw(rr *proto.RouteRaw, peerIDByIndex []string) *nbroute.Route {
|
||||
return r
|
||||
}
|
||||
|
||||
func decodeNameServerGroupRaw(nsg *proto.NameServerGroupRaw) *nbdns.NameServerGroup {
|
||||
out := &nbdns.NameServerGroup{
|
||||
func decodeNameServerGroupRaw(nsg *proto.NameServerGroupRaw) *nmdata.NameServerGroup {
|
||||
out := &nmdata.NameServerGroup{
|
||||
ID: nsg.Id,
|
||||
PublicID: nsg.Id,
|
||||
Groups: nsg.GroupIds,
|
||||
@@ -388,13 +478,13 @@ func decodeNameServerGroupRaw(nsg *proto.NameServerGroupRaw) *nbdns.NameServerGr
|
||||
Domains: nsg.Domains,
|
||||
Enabled: nsg.Enabled,
|
||||
SearchDomainsEnabled: nsg.SearchDomainsEnabled,
|
||||
NameServers: make([]nbdns.NameServer, 0, len(nsg.Nameservers)),
|
||||
NameServers: make([]nmdata.NameServer, 0, len(nsg.Nameservers)),
|
||||
}
|
||||
for _, ns := range nsg.Nameservers {
|
||||
if addr, err := netip.ParseAddr(ns.IP); err == nil {
|
||||
out.NameServers = append(out.NameServers, nbdns.NameServer{
|
||||
out.NameServers = append(out.NameServers, nmdata.NameServer{
|
||||
IP: addr,
|
||||
NSType: nbdns.NameServerType(ns.NSType),
|
||||
NSType: int(ns.NSType),
|
||||
Port: int(ns.Port),
|
||||
})
|
||||
}
|
||||
@@ -402,14 +492,14 @@ func decodeNameServerGroupRaw(nsg *proto.NameServerGroupRaw) *nbdns.NameServerGr
|
||||
return out
|
||||
}
|
||||
|
||||
func decodeNetworkResource(nr *proto.NetworkResourceRaw) *types.ComponentResource {
|
||||
out := &types.ComponentResource{
|
||||
func decodeNetworkResource(nr *proto.NetworkResourceRaw) *nmdata.NetworkResource {
|
||||
out := &nmdata.NetworkResource{
|
||||
ID: nr.Id,
|
||||
PublicID: nr.Id,
|
||||
NetworkID: nr.NetworkSeq,
|
||||
Name: nr.Name,
|
||||
Description: nr.Description,
|
||||
Type: types.ComponentResourceType(nr.Type),
|
||||
Type: nr.Type,
|
||||
Address: nr.Address,
|
||||
Domain: nr.DomainValue,
|
||||
Enabled: nr.Enabled,
|
||||
@@ -422,10 +512,10 @@ func decodeNetworkResource(nr *proto.NetworkResourceRaw) *types.ComponentResourc
|
||||
return out
|
||||
}
|
||||
|
||||
func decodeSimpleRecords(records []*proto.SimpleRecord) []nbdns.SimpleRecord {
|
||||
out := make([]nbdns.SimpleRecord, 0, len(records))
|
||||
func decodeSimpleRecords(records []*proto.SimpleRecord) []nmdata.SimpleRecord {
|
||||
out := make([]nmdata.SimpleRecord, 0, len(records))
|
||||
for _, r := range records {
|
||||
out = append(out, nbdns.SimpleRecord{
|
||||
out = append(out, nmdata.SimpleRecord{
|
||||
Name: r.Name,
|
||||
Type: int(r.Type),
|
||||
Class: r.Class,
|
||||
@@ -436,10 +526,10 @@ func decodeSimpleRecords(records []*proto.SimpleRecord) []nbdns.SimpleRecord {
|
||||
return out
|
||||
}
|
||||
|
||||
func decodeCustomZones(zones []*proto.CustomZone) []nbdns.CustomZone {
|
||||
out := make([]nbdns.CustomZone, 0, len(zones))
|
||||
func decodeCustomZones(zones []*proto.CustomZone) []nmdata.CustomZone {
|
||||
out := make([]nmdata.CustomZone, 0, len(zones))
|
||||
for _, z := range zones {
|
||||
out = append(out, nbdns.CustomZone{
|
||||
out = append(out, nmdata.CustomZone{
|
||||
Domain: z.Domain,
|
||||
Records: decodeSimpleRecords(z.Records),
|
||||
SearchDomainDisabled: z.SearchDomainDisabled,
|
||||
@@ -460,16 +550,16 @@ func uint32SliceToStrings(ports []uint32) []string {
|
||||
return out
|
||||
}
|
||||
|
||||
func portRangesFromProto(ranges []*proto.PortInfo_Range) []types.RulePortRange {
|
||||
func portRangesFromProto(ranges []*proto.PortInfo_Range) []nmdata.RulePortRange {
|
||||
if len(ranges) == 0 {
|
||||
return nil
|
||||
}
|
||||
out := make([]types.RulePortRange, 0, len(ranges))
|
||||
out := make([]nmdata.RulePortRange, 0, len(ranges))
|
||||
for _, r := range ranges {
|
||||
if r == nil || r.Start > 65535 || r.End > 65535 {
|
||||
continue
|
||||
}
|
||||
out = append(out, types.RulePortRange{
|
||||
out = append(out, nmdata.RulePortRange{
|
||||
Start: uint16(r.Start),
|
||||
End: uint16(r.End),
|
||||
})
|
||||
|
||||
40
shared/management/networkmap/decode_test.go
Normal file
40
shared/management/networkmap/decode_test.go
Normal file
@@ -0,0 +1,40 @@
|
||||
package networkmap
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestDecodePolicy(t *testing.T) {
|
||||
assert.Equal(t,
|
||||
resourceFromProto(
|
||||
&proto.ResourceCompact{Type: proto.ResourceCompactType_peer, ResourceId: &proto.ResourceCompact_PeerIndex{PeerIndex: uint32(1)}},
|
||||
[]string{"invalid-id-0", "valid-id", "invalid-id-2"}),
|
||||
nmdata.Resource{Type: "peer", ID: "valid-id"})
|
||||
// check invalid peer index returns an empty resource
|
||||
assert.Equal(t,
|
||||
resourceFromProto(
|
||||
&proto.ResourceCompact{Type: proto.ResourceCompactType_peer, ResourceId: &proto.ResourceCompact_PeerIndex{PeerIndex: uint32(100)}},
|
||||
[]string{"invalid-id-0", "valid-id", "invalid-id-2"}),
|
||||
nmdata.Resource{})
|
||||
assert.Equal(t,
|
||||
resourceFromProto(
|
||||
&proto.ResourceCompact{Type: proto.ResourceCompactType_domain, ResourceId: &proto.ResourceCompact_Id{Id: "domain"}}, []string{}),
|
||||
nmdata.Resource{Type: "domain", ID: "domain"})
|
||||
assert.Equal(t,
|
||||
resourceFromProto(
|
||||
&proto.ResourceCompact{Type: proto.ResourceCompactType_host, ResourceId: &proto.ResourceCompact_Id{Id: "host"}}, []string{}),
|
||||
nmdata.Resource{Type: "host", ID: "host"})
|
||||
assert.Equal(t,
|
||||
resourceFromProto(
|
||||
&proto.ResourceCompact{Type: proto.ResourceCompactType_subnet, ResourceId: &proto.ResourceCompact_Id{Id: "subnet"}}, []string{}),
|
||||
nmdata.Resource{Type: "subnet", ID: "subnet"})
|
||||
// an unknown resource type return an empty resource
|
||||
assert.Equal(t,
|
||||
resourceFromProto(
|
||||
&proto.ResourceCompact{Type: proto.ResourceCompactType_unknown_type, ResourceId: &proto.ResourceCompact_Id{Id: "boom"}}, []string{}),
|
||||
nmdata.Resource{})
|
||||
}
|
||||
@@ -20,7 +20,7 @@ import (
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
"net/netip"
|
||||
|
||||
nbroute "github.com/netbirdio/netbird/route"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
"github.com/netbirdio/netbird/shared/management/types"
|
||||
"github.com/netbirdio/netbird/shared/netiputil"
|
||||
@@ -28,7 +28,7 @@ import (
|
||||
)
|
||||
|
||||
// ToProtocolRoutes converts a slice of typed routes to their proto form.
|
||||
func ToProtocolRoutes(routes []*nbroute.Route) []*proto.Route {
|
||||
func ToProtocolRoutes(routes []*nmdata.Route) []*proto.Route {
|
||||
protoRoutes := make([]*proto.Route, 0, len(routes))
|
||||
for _, r := range routes {
|
||||
protoRoutes = append(protoRoutes, ToProtocolRoute(r))
|
||||
@@ -37,7 +37,7 @@ func ToProtocolRoutes(routes []*nbroute.Route) []*proto.Route {
|
||||
}
|
||||
|
||||
// ToProtocolRoute converts one typed route to its proto form.
|
||||
func ToProtocolRoute(route *nbroute.Route) *proto.Route {
|
||||
func ToProtocolRoute(route *nmdata.Route) *proto.Route {
|
||||
return &proto.Route{
|
||||
ID: string(route.ID),
|
||||
NetID: string(route.NetID),
|
||||
@@ -273,7 +273,7 @@ func ToProtocolDNSConfig(update nbdns.Config, cache DNSConfigCache, forwardPort
|
||||
|
||||
// AppendRemotePeerConfig appends typed peers as proto.RemotePeerConfig
|
||||
// entries to dst and returns the result.
|
||||
func AppendRemotePeerConfig(dst []*proto.RemotePeerConfig, peers []*types.ComponentPeer, dnsName string, includeIPv6 bool) []*proto.RemotePeerConfig {
|
||||
func AppendRemotePeerConfig(dst []*proto.RemotePeerConfig, peers []*nmdata.Peer, dnsName string, includeIPv6 bool) []*proto.RemotePeerConfig {
|
||||
for _, rPeer := range peers {
|
||||
allowedIPs := []string{rPeer.IP.String() + "/32"}
|
||||
if includeIPv6 && rPeer.IPv6.IsValid() {
|
||||
@@ -284,7 +284,7 @@ func AppendRemotePeerConfig(dst []*proto.RemotePeerConfig, peers []*types.Compon
|
||||
AllowedIps: allowedIPs,
|
||||
SshConfig: &proto.SSHConfig{SshPubKey: []byte(rPeer.SSHKey)},
|
||||
Fqdn: rPeer.FQDN(dnsName),
|
||||
AgentVersion: rPeer.AgentVersion,
|
||||
AgentVersion: rPeer.Meta.WtVersion,
|
||||
})
|
||||
}
|
||||
return dst
|
||||
|
||||
@@ -54,8 +54,8 @@ func EnvelopeToNetworkMap(ctx context.Context, env *proto.NetworkMapEnvelope, lo
|
||||
}
|
||||
components.PeerID = canonicalKey
|
||||
|
||||
includeIPv6 := localPeer.SupportsIPv6 && localPeer.IPv6.IsValid()
|
||||
useSourcePrefixes := localPeer.SupportsSourcePrefixes
|
||||
includeIPv6 := localPeer.SupportsIPv6() && localPeer.IPv6.IsValid()
|
||||
useSourcePrefixes := localPeer.SupportsSourcePrefixes()
|
||||
|
||||
typedNM := components.Calculate(ctx)
|
||||
|
||||
|
||||
@@ -15,6 +15,7 @@ import (
|
||||
mgmtgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
nbnetworkmap "github.com/netbirdio/netbird/shared/management/networkmap"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
@@ -55,13 +56,13 @@ func TestEnvelopeToNetworkMap_RoundTrip(t *testing.T) {
|
||||
func TestCalculate_FirewallRuleProtocol_NeverNetbirdSSH(t *testing.T) {
|
||||
c, localPeerKey := buildSmokeComponents(t)
|
||||
// Replace the smoke policy with a NetbirdSSH-protocol allow.
|
||||
c.Policies = []*types.Policy{{
|
||||
c.Policies = []*nmdata.Policy{{
|
||||
ID: "pol-ssh", PublicID: "2", Enabled: true,
|
||||
Rules: []*types.PolicyRule{{
|
||||
Rules: []*nmdata.PolicyRule{{
|
||||
ID: "rule-ssh",
|
||||
Enabled: true,
|
||||
Action: types.PolicyTrafficActionAccept,
|
||||
Protocol: types.PolicyRuleProtocolNetbirdSSH,
|
||||
Action: string(types.PolicyTrafficActionAccept),
|
||||
Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
|
||||
Bidirectional: true,
|
||||
Sources: []string{"group-all"},
|
||||
Destinations: []string{"group-all"},
|
||||
@@ -143,39 +144,39 @@ func TestDecodeEnvelope_MalformedWgKeyPeerSkipped(t *testing.T) {
|
||||
func TestEnvelopeRoundTrip_AllGroupShortCircuitParity(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
peers := map[string]*types.ComponentPeer{}
|
||||
peers := map[string]*nmdata.Peer{}
|
||||
for i, id := range []string{"peer-T", "peer-S", "peer-ALL", "peer-O"} {
|
||||
peers[id] = &types.ComponentPeer{
|
||||
ID: id,
|
||||
Key: randomWgKey(t),
|
||||
IP: netip.AddrFrom4([4]byte{100, 64, 0, byte(i + 1)}),
|
||||
DNSLabel: id,
|
||||
AgentVersion: "0.40.0",
|
||||
peers[id] = &nmdata.Peer{
|
||||
ID: id,
|
||||
Key: randomWgKey(t),
|
||||
IP: netip.AddrFrom4([4]byte{100, 64, 0, byte(i + 1)}),
|
||||
DNSLabel: id,
|
||||
Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
|
||||
}
|
||||
}
|
||||
|
||||
c := &types.NetworkMapComponents{
|
||||
PeerID: "peer-T",
|
||||
Network: &types.Network{
|
||||
Network: &nmdata.Network{
|
||||
Identifier: "net-all-groups",
|
||||
Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)},
|
||||
Serial: 1,
|
||||
},
|
||||
AccountSettings: &types.AccountSettingsInfo{},
|
||||
DNSSettings: &types.DNSSettings{},
|
||||
AccountSettings: &nmdata.AccountSettingsInfo{},
|
||||
DNSSettings: &nmdata.DNSSettings{},
|
||||
Peers: peers,
|
||||
Groups: map[string]*types.ComponentGroup{
|
||||
"g-src": {ID: "g-src", PublicID: "1", Name: "staff", Peers: []string{"peer-T", "peer-S"}},
|
||||
"g-all": {ID: "g-all", PublicID: "2", Name: "All", Peers: []string{"peer-ALL"}},
|
||||
"g-two": {ID: "g-two", PublicID: "3", Name: "second", Peers: []string{"peer-T", "peer-O"}},
|
||||
Groups: map[string]*nmdata.Group{
|
||||
"g-src": {PublicID: "1", Name: "staff", Peers: []string{"peer-T", "peer-S"}},
|
||||
"g-all": {PublicID: "2", Name: "All", Peers: []string{"peer-ALL"}},
|
||||
"g-two": {PublicID: "3", Name: "second", Peers: []string{"peer-T", "peer-O"}},
|
||||
},
|
||||
Policies: []*types.Policy{{
|
||||
Policies: []*nmdata.Policy{{
|
||||
ID: "pol-multi-dest", PublicID: "10", Enabled: true,
|
||||
Rules: []*types.PolicyRule{{
|
||||
Rules: []*nmdata.PolicyRule{{
|
||||
ID: "rule-multi-dest",
|
||||
Enabled: true,
|
||||
Action: types.PolicyTrafficActionAccept,
|
||||
Protocol: types.PolicyRuleProtocolALL,
|
||||
Action: string(types.PolicyTrafficActionAccept),
|
||||
Protocol: string(types.PolicyRuleProtocolALL),
|
||||
Sources: []string{"g-src"},
|
||||
Destinations: []string{"g-all", "g-two"},
|
||||
}},
|
||||
@@ -231,33 +232,33 @@ func buildSmokeComponents(t *testing.T) (*types.NetworkMapComponents, string) {
|
||||
peerAKey := randomWgKey(t)
|
||||
peerBKey := randomWgKey(t)
|
||||
|
||||
peerA := &types.ComponentPeer{
|
||||
ID: "peer-A",
|
||||
Key: peerAKey,
|
||||
IP: netip.AddrFrom4([4]byte{100, 64, 0, 1}),
|
||||
DNSLabel: "peerA",
|
||||
AgentVersion: "0.40.0",
|
||||
peerA := &nmdata.Peer{
|
||||
ID: "peer-A",
|
||||
Key: peerAKey,
|
||||
IP: netip.AddrFrom4([4]byte{100, 64, 0, 1}),
|
||||
DNSLabel: "peerA",
|
||||
Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
|
||||
}
|
||||
peerB := &types.ComponentPeer{
|
||||
ID: "peer-B",
|
||||
Key: peerBKey,
|
||||
IP: netip.AddrFrom4([4]byte{100, 64, 0, 2}),
|
||||
DNSLabel: "peerB",
|
||||
AgentVersion: "0.40.0",
|
||||
peerB := &nmdata.Peer{
|
||||
ID: "peer-B",
|
||||
Key: peerBKey,
|
||||
IP: netip.AddrFrom4([4]byte{100, 64, 0, 2}),
|
||||
DNSLabel: "peerB",
|
||||
Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
|
||||
}
|
||||
|
||||
group := &types.ComponentGroup{
|
||||
ID: "group-all", PublicID: "1", Name: "All",
|
||||
group := &nmdata.Group{
|
||||
PublicID: "1", Name: "All",
|
||||
Peers: []string{"peer-A", "peer-B"},
|
||||
}
|
||||
|
||||
policy := &types.Policy{
|
||||
policy := &nmdata.Policy{
|
||||
ID: "pol-allow", PublicID: "1", Enabled: true,
|
||||
Rules: []*types.PolicyRule{{
|
||||
Rules: []*nmdata.PolicyRule{{
|
||||
ID: "rule-allow",
|
||||
Enabled: true,
|
||||
Action: types.PolicyTrafficActionAccept,
|
||||
Protocol: types.PolicyRuleProtocolALL,
|
||||
Action: string(types.PolicyTrafficActionAccept),
|
||||
Protocol: string(types.PolicyRuleProtocolALL),
|
||||
Bidirectional: true,
|
||||
Sources: []string{"group-all"},
|
||||
Destinations: []string{"group-all"},
|
||||
@@ -266,21 +267,21 @@ func buildSmokeComponents(t *testing.T) (*types.NetworkMapComponents, string) {
|
||||
|
||||
c := &types.NetworkMapComponents{
|
||||
PeerID: "peer-A",
|
||||
Network: &types.Network{
|
||||
Network: &nmdata.Network{
|
||||
Identifier: "net-smoke",
|
||||
Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)},
|
||||
Serial: 1,
|
||||
},
|
||||
AccountSettings: &types.AccountSettingsInfo{},
|
||||
DNSSettings: &types.DNSSettings{},
|
||||
Peers: map[string]*types.ComponentPeer{
|
||||
AccountSettings: &nmdata.AccountSettingsInfo{},
|
||||
DNSSettings: &nmdata.DNSSettings{},
|
||||
Peers: map[string]*nmdata.Peer{
|
||||
"peer-A": peerA,
|
||||
"peer-B": peerB,
|
||||
},
|
||||
Groups: map[string]*types.ComponentGroup{
|
||||
Groups: map[string]*nmdata.Group{
|
||||
"group-all": group,
|
||||
},
|
||||
Policies: []*types.Policy{policy},
|
||||
Policies: []*nmdata.Policy{policy},
|
||||
}
|
||||
return c, peerAKey
|
||||
}
|
||||
|
||||
659
shared/management/networkmap/networkmapcompute.go
Normal file
659
shared/management/networkmap/networkmapcompute.go
Normal file
@@ -0,0 +1,659 @@
|
||||
package networkmap
|
||||
|
||||
import (
|
||||
"slices"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
"github.com/netbirdio/netbird/shared/management/types"
|
||||
)
|
||||
|
||||
type sshRequirements struct {
|
||||
neededGroupIDs map[string]struct{}
|
||||
needAllowedUserIDs bool
|
||||
}
|
||||
|
||||
// GetPeerNetworkMapComponents computes the peer's NetworkMapComponents from the
|
||||
// slim twin store. It mirrors the former Account.GetPeerNetworkMapComponents
|
||||
// exactly, operating on nmdata twins throughout — no Account reference and no
|
||||
// twin↔real conversion, since the produced components hold twins.
|
||||
func (nmd *NetworkMapData) GetPeerNetworkMapComponents(peerID string, peersCustomZone nmdata.CustomZone) *types.NetworkMapComponents {
|
||||
peer := nmd.Peers[peerID]
|
||||
if peer == nil {
|
||||
return types.EmptyNetworkMapComponents(&types.NetworkMapComponents{
|
||||
PeerID: peerID,
|
||||
Network: nmd.Network,
|
||||
Peers: map[string]*nmdata.Peer{peerID: peer},
|
||||
})
|
||||
}
|
||||
|
||||
if _, ok := nmd.ValidatedPeers[peerID]; !ok {
|
||||
return types.EmptyNetworkMapComponents(&types.NetworkMapComponents{
|
||||
PeerID: peerID,
|
||||
Network: nmd.Network,
|
||||
Peers: map[string]*nmdata.Peer{peerID: peer},
|
||||
})
|
||||
}
|
||||
|
||||
components := &types.NetworkMapComponents{
|
||||
PeerID: peerID,
|
||||
Network: nmd.Network,
|
||||
AccountSettings: nmd.AccountSettings,
|
||||
DNSSettings: nmd.DNSSettings,
|
||||
CustomZoneDomain: peersCustomZone.Domain,
|
||||
NameServerGroups: make([]*nmdata.NameServerGroup, 0),
|
||||
ResourcePoliciesMap: make(map[string][]*nmdata.Policy),
|
||||
RoutersMap: make(map[string]map[string]*nmdata.NetworkRouter),
|
||||
NetworkResources: make([]*nmdata.NetworkResource, 0),
|
||||
PostureFailedPeers: make(map[string]map[string]struct{}, len(nmd.PostureChecks)),
|
||||
RouterPeers: make(map[string]*nmdata.Peer),
|
||||
NetworkXIDToPublicID: nmd.NetworkXIDToPublicID,
|
||||
PostureCheckXIDToPublicID: nmd.PostureCheckXIDToPublicID,
|
||||
}
|
||||
|
||||
relevantPeers, relevantGroups, relevantPolicies, relevantRoutes, sshReqs := nmd.getPeersGroupsPoliciesRoutes(peerID, peer.SSHEnabled, &components.PostureFailedPeers)
|
||||
|
||||
if len(sshReqs.neededGroupIDs) > 0 {
|
||||
components.GroupIDToUserIDs = filterGroupIDToUserIDs(nmd.GroupIDToUserIDs, sshReqs.neededGroupIDs)
|
||||
}
|
||||
if sshReqs.needAllowedUserIDs {
|
||||
components.AllowedUserIDs = nmd.getAllowedUserIDs()
|
||||
}
|
||||
|
||||
components.Peers = relevantPeers
|
||||
components.Groups = relevantGroups
|
||||
components.Policies = relevantPolicies
|
||||
components.Routes = relevantRoutes
|
||||
components.AllDNSRecords = filterDNSRecordsByPeers(peersCustomZone.Records, relevantPeers, peer.SupportsIPv6() && peer.IPv6.IsValid())
|
||||
|
||||
peerGroups := nmd.GetPeerGroups(peerID)
|
||||
components.AccountZones = nmd.appliedZones(peerGroups)
|
||||
components.AccountZones = append(components.AccountZones, nmd.privateServiceZones(peerGroups)...)
|
||||
|
||||
for _, nsGroup := range nmd.NameServerGroups {
|
||||
if nsGroup.Enabled {
|
||||
for _, gID := range nsGroup.Groups {
|
||||
if _, found := relevantGroups[gID]; found {
|
||||
components.NameServerGroups = append(components.NameServerGroups, nsGroup)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for _, resource := range nmd.NetworkResources {
|
||||
if !resource.Enabled {
|
||||
continue
|
||||
}
|
||||
|
||||
policies, exists := nmd.ResourcePolicies[resource.ID]
|
||||
if !exists {
|
||||
continue
|
||||
}
|
||||
|
||||
addSourcePeers := false
|
||||
|
||||
networkRoutingPeers, routerExists := nmd.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 == string(types.ResourceTypePeer) && policy.Rules[0].SourceResource.ID != "" {
|
||||
peers = []string{policy.Rules[0].SourceResource.ID}
|
||||
} else {
|
||||
peers = nmd.getUniquePeerIDsFromGroupsIDs(policy.SourceGroups())
|
||||
}
|
||||
for _, pID := range nmd.getPostureValidPeersSaveFailed(peers, policy.SourcePostureChecks, &components.PostureFailedPeers) {
|
||||
if _, exists := components.Peers[pID]; !exists {
|
||||
components.Peers[pID] = nmd.Peers[pID]
|
||||
}
|
||||
}
|
||||
} else {
|
||||
peerInSources := false
|
||||
if policy.Rules[0].SourceResource.Type == string(types.ResourceTypePeer) && policy.Rules[0].SourceResource.ID != "" {
|
||||
peerInSources = policy.Rules[0].SourceResource.ID == peerID
|
||||
} else {
|
||||
for _, groupID := range policy.SourceGroups() {
|
||||
if group := nmd.Groups[groupID]; group != nil && slices.Contains(group.Peers, peerID) {
|
||||
peerInSources = true
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
if !peerInSources {
|
||||
continue
|
||||
}
|
||||
isValid, pname := nmd.validatePostureChecksOnPeerGetFailed(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 := nmd.Groups[srcGroupID]; g != nil {
|
||||
if _, exists := components.Groups[srcGroupID]; !exists {
|
||||
components.Groups[srcGroupID] = g
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, dstGroupID := range rule.Destinations {
|
||||
if g := nmd.Groups[dstGroupID]; g != nil {
|
||||
if _, exists := components.Groups[dstGroupID]; !exists {
|
||||
components.Groups[dstGroupID] = g
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
components.ResourcePoliciesMap[resource.ID] = policies
|
||||
}
|
||||
|
||||
if addSourcePeers {
|
||||
components.RoutersMap[resource.NetworkID] = networkRoutingPeers
|
||||
for peerIDKey := range networkRoutingPeers {
|
||||
if p := nmd.Peers[peerIDKey]; p != nil {
|
||||
if _, exists := components.RouterPeers[peerIDKey]; !exists {
|
||||
components.RouterPeers[peerIDKey] = p
|
||||
}
|
||||
if _, exists := components.Peers[peerIDKey]; !exists {
|
||||
if _, validated := nmd.ValidatedPeers[peerIDKey]; validated {
|
||||
components.Peers[peerIDKey] = p
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
components.NetworkResources = append(components.NetworkResources, resource)
|
||||
}
|
||||
}
|
||||
|
||||
filterGroupPeers(&components.Groups, components.Peers)
|
||||
filterPostureFailedPeers(&components.PostureFailedPeers, components.Policies, components.ResourcePoliciesMap, components.Peers)
|
||||
|
||||
return components
|
||||
}
|
||||
|
||||
func (nmd *NetworkMapData) getPeersGroupsPoliciesRoutes(
|
||||
peerID string,
|
||||
peerSSHEnabled bool,
|
||||
postureFailedPeers *map[string]map[string]struct{},
|
||||
) (map[string]*nmdata.Peer, map[string]*nmdata.Group, []*nmdata.Policy, []*nmdata.Route, sshRequirements) {
|
||||
relevantPeerIDs := make(map[string]*nmdata.Peer, len(nmd.Peers)/4)
|
||||
relevantGroupIDs := make(map[string]*nmdata.Group, len(nmd.Groups)/4)
|
||||
relevantPolicies := make([]*nmdata.Policy, 0, len(nmd.Policies))
|
||||
relevantRoutes := make([]*nmdata.Route, 0, len(nmd.Routes))
|
||||
sshReqs := sshRequirements{neededGroupIDs: make(map[string]struct{})}
|
||||
|
||||
relevantPeerIDs[peerID] = nmd.Peers[peerID]
|
||||
|
||||
peerGroupSet := make(map[string]struct{}, 8)
|
||||
for groupID, group := range nmd.Groups {
|
||||
if slices.Contains(group.Peers, peerID) {
|
||||
relevantGroupIDs[groupID] = group
|
||||
peerGroupSet[groupID] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
routeAccessControlGroups := make(map[string]struct{})
|
||||
for _, r := range nmd.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] = nmd.Groups[groupID]
|
||||
}
|
||||
for _, groupID := range r.Groups {
|
||||
relevantGroupIDs[groupID] = nmd.Groups[groupID]
|
||||
}
|
||||
if r.Enabled {
|
||||
for _, groupID := range r.AccessControlGroups {
|
||||
relevantGroupIDs[groupID] = nmd.Groups[groupID]
|
||||
routeAccessControlGroups[groupID] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
if r.Peer != "" {
|
||||
if _, ok := nmd.ValidatedPeers[r.Peer]; ok {
|
||||
if p := nmd.Peers[r.Peer]; p != nil {
|
||||
relevantPeerIDs[r.Peer] = p
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, groupID := range r.PeerGroups {
|
||||
g := nmd.Groups[groupID]
|
||||
if g == nil {
|
||||
continue
|
||||
}
|
||||
for _, pid := range g.Peers {
|
||||
if _, exists := relevantPeerIDs[pid]; exists {
|
||||
continue
|
||||
}
|
||||
if _, ok := nmd.ValidatedPeers[pid]; !ok {
|
||||
continue
|
||||
}
|
||||
if p := nmd.Peers[pid]; p != nil {
|
||||
relevantPeerIDs[pid] = p
|
||||
}
|
||||
}
|
||||
}
|
||||
relevantRoutes = append(relevantRoutes, r)
|
||||
}
|
||||
|
||||
for _, policy := range nmd.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] = nmd.Groups[srcGroupID]
|
||||
}
|
||||
for _, dstGroupID := range rule.Destinations {
|
||||
relevantGroupIDs[dstGroupID] = nmd.Groups[dstGroupID]
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var sourcePeers, destinationPeers []string
|
||||
var peerInSources, peerInDestinations bool
|
||||
|
||||
if rule.SourceResource.Type == string(types.ResourceTypePeer) && rule.SourceResource.ID != "" {
|
||||
sourcePeers = []string{rule.SourceResource.ID}
|
||||
if rule.SourceResource.ID == peerID {
|
||||
peerInSources = true
|
||||
}
|
||||
} else {
|
||||
sourcePeers, peerInSources = nmd.getPeersFromGroups(rule.Sources, peerID, policy.SourcePostureChecks, postureFailedPeers)
|
||||
}
|
||||
|
||||
if rule.DestinationResource.Type == string(types.ResourceTypePeer) && rule.DestinationResource.ID != "" {
|
||||
destinationPeers = []string{rule.DestinationResource.ID}
|
||||
if rule.DestinationResource.ID == peerID {
|
||||
peerInDestinations = true
|
||||
}
|
||||
} else {
|
||||
destinationPeers, peerInDestinations = nmd.getPeersFromGroups(rule.Destinations, peerID, nil, postureFailedPeers)
|
||||
}
|
||||
|
||||
if peerInSources {
|
||||
policyRelevant = true
|
||||
for _, pid := range destinationPeers {
|
||||
relevantPeerIDs[pid] = nmd.Peers[pid]
|
||||
}
|
||||
for _, dstGroupID := range rule.Destinations {
|
||||
relevantGroupIDs[dstGroupID] = nmd.Groups[dstGroupID]
|
||||
}
|
||||
}
|
||||
|
||||
if peerInDestinations {
|
||||
policyRelevant = true
|
||||
for _, pid := range sourcePeers {
|
||||
relevantPeerIDs[pid] = nmd.Peers[pid]
|
||||
}
|
||||
for _, srcGroupID := range rule.Sources {
|
||||
relevantGroupIDs[srcGroupID] = nmd.Groups[srcGroupID]
|
||||
}
|
||||
|
||||
if rule.Protocol == string(types.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 nmdata.PolicyRuleImpliesLegacySSH(rule) && peerSSHEnabled {
|
||||
sshReqs.needAllowedUserIDs = true
|
||||
}
|
||||
}
|
||||
}
|
||||
if policyRelevant {
|
||||
relevantPolicies = append(relevantPolicies, policy)
|
||||
}
|
||||
}
|
||||
|
||||
return relevantPeerIDs, relevantGroupIDs, relevantPolicies, relevantRoutes, sshReqs
|
||||
}
|
||||
|
||||
func (nmd *NetworkMapData) getPeersFromGroups(groups []string, peerID string, sourcePostureChecksIDs []string,
|
||||
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 := nmd.Groups[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 := nmd.Peers[pid]
|
||||
if !ok || peer == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
if _, ok := nmd.ValidatedPeers[peer.ID]; !ok {
|
||||
continue
|
||||
}
|
||||
|
||||
isValid, pname := nmd.validatePostureChecksOnPeerGetFailed(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 := nmd.Peers[pid]
|
||||
if !ok || peer == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
if _, ok := nmd.ValidatedPeers[peer.ID]; !ok {
|
||||
continue
|
||||
}
|
||||
|
||||
isValid, pname := nmd.validatePostureChecksOnPeerGetFailed(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 (nmd *NetworkMapData) validatePostureChecksOnPeerGetFailed(sourcePostureChecksID []string, peerID string) (bool, string) {
|
||||
peer, ok := nmd.Peers[peerID]
|
||||
if !ok || peer == nil {
|
||||
return false, ""
|
||||
}
|
||||
|
||||
for _, postureChecksID := range sourcePostureChecksID {
|
||||
postureChecks := nmd.PostureChecks[postureChecksID]
|
||||
if postureChecks == nil {
|
||||
continue
|
||||
}
|
||||
if !postureChecks.Passes(peer) {
|
||||
return false, postureChecksID
|
||||
}
|
||||
}
|
||||
return true, ""
|
||||
}
|
||||
|
||||
func (nmd *NetworkMapData) getPostureValidPeersSaveFailed(inputPeers []string, postureChecksIDs []string, postureFailedPeers *map[string]map[string]struct{}) []string {
|
||||
var dest []string
|
||||
for _, peerID := range inputPeers {
|
||||
if _, validated := nmd.ValidatedPeers[peerID]; !validated {
|
||||
continue
|
||||
}
|
||||
valid, pname := nmd.validatePostureChecksOnPeerGetFailed(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
|
||||
}
|
||||
|
||||
func (nmd *NetworkMapData) GetPeerGroups(peerID string) map[string]struct{} {
|
||||
groups := make(map[string]struct{})
|
||||
for groupID, group := range nmd.Groups {
|
||||
if slices.Contains(group.Peers, peerID) {
|
||||
groups[groupID] = struct{}{}
|
||||
}
|
||||
}
|
||||
return groups
|
||||
}
|
||||
|
||||
func (nmd *NetworkMapData) getUniquePeerIDsFromGroupsIDs(groups []string) []string {
|
||||
peerIDs := make(map[string]struct{}, len(groups))
|
||||
for _, groupID := range groups {
|
||||
group := nmd.Groups[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 (nmd *NetworkMapData) getAllowedUserIDs() map[string]struct{} {
|
||||
return nmd.AllowedUserIDs
|
||||
}
|
||||
|
||||
func (nmd *NetworkMapData) appliedZones(peerGroups map[string]struct{}) []nmdata.CustomZone {
|
||||
if len(peerGroups) == 0 {
|
||||
return nil
|
||||
}
|
||||
var out []nmdata.CustomZone
|
||||
for _, cand := range nmd.AppliedZoneCandidates {
|
||||
if peerInDistributionGroups(peerGroups, cand.DistributionGroups) {
|
||||
out = append(out, cand.Zone)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func (nmd *NetworkMapData) privateServiceZones(peerGroups map[string]struct{}) []nmdata.CustomZone {
|
||||
byApex := make(map[string]*nmdata.CustomZone)
|
||||
var order []string
|
||||
for _, cand := range nmd.PrivateServiceCandidates {
|
||||
if !peerInDistributionGroups(peerGroups, cand.AccessGroups) {
|
||||
continue
|
||||
}
|
||||
zone, exists := byApex[cand.Zone.Domain]
|
||||
if !exists {
|
||||
nz := nmdata.CustomZone{
|
||||
Domain: cand.Zone.Domain,
|
||||
SearchDomainDisabled: cand.Zone.SearchDomainDisabled,
|
||||
NonAuthoritative: cand.Zone.NonAuthoritative,
|
||||
}
|
||||
byApex[cand.Zone.Domain] = &nz
|
||||
zone = &nz
|
||||
order = append(order, cand.Zone.Domain)
|
||||
}
|
||||
zone.Records = append(zone.Records, cand.Zone.Records...)
|
||||
}
|
||||
|
||||
var out []nmdata.CustomZone
|
||||
for _, apex := range order {
|
||||
zone := byApex[apex]
|
||||
if len(zone.Records) == 0 {
|
||||
continue
|
||||
}
|
||||
out = append(out, *zone)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func peerInDistributionGroups(peerGroups map[string]struct{}, groups []string) bool {
|
||||
for _, g := range groups {
|
||||
if _, ok := peerGroups[g]; ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func filterGroupPeers(groups *map[string]*nmdata.Group, peers map[string]*nmdata.Peer) {
|
||||
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.Copy()
|
||||
ng.Peers = filteredPeers
|
||||
(*groups)[groupID] = ng
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func filterPostureFailedPeers(postureFailedPeers *map[string]map[string]struct{}, policies []*nmdata.Policy, resourcePoliciesMap map[string][]*nmdata.Policy, peers map[string]*nmdata.Peer) {
|
||||
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 []nmdata.SimpleRecord, peers map[string]*nmdata.Peer, includeIPv6 bool) []nmdata.SimpleRecord {
|
||||
if len(records) == 0 || len(peers) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
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([]nmdata.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
|
||||
}
|
||||
55
shared/management/networkmap/networkmapdata.go
Normal file
55
shared/management/networkmap/networkmapdata.go
Normal file
@@ -0,0 +1,55 @@
|
||||
package networkmap
|
||||
|
||||
import (
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
|
||||
)
|
||||
|
||||
// NetworkMapData is a dependency-light, slim twin of the server Account. It
|
||||
// carries only the state GetPeerNetworkMapComponents needs, expressed in the
|
||||
// fresh nmdata twin types. A builder converts an Account into a NetworkMapData
|
||||
// once per account; the per-peer components calculation then runs on this twin
|
||||
// with no reference back to the Account.
|
||||
type NetworkMapData struct {
|
||||
Peers map[string]*nmdata.Peer
|
||||
Groups map[string]*nmdata.Group
|
||||
Policies []*nmdata.Policy
|
||||
Routes []*nmdata.Route
|
||||
NameServerGroups []*nmdata.NameServerGroup
|
||||
NetworkResources []*nmdata.NetworkResource
|
||||
|
||||
Network *nmdata.Network
|
||||
DNSSettings *nmdata.DNSSettings
|
||||
AccountSettings *nmdata.AccountSettingsInfo
|
||||
|
||||
PostureChecks map[string]*nmdata.PostureChecks
|
||||
|
||||
AllowedUserIDs map[string]struct{}
|
||||
NetworkXIDToPublicID map[string]string
|
||||
PostureCheckXIDToPublicID map[string]string
|
||||
ValidatedPeers map[string]struct{}
|
||||
ResourcePolicies map[string][]*nmdata.Policy
|
||||
Routers map[string]map[string]*nmdata.NetworkRouter
|
||||
GroupIDToUserIDs map[string][]string
|
||||
DNSDomain string
|
||||
|
||||
AppliedZoneCandidates []AppliedZoneCandidate
|
||||
PrivateServiceCandidates []PrivateServiceCandidate
|
||||
}
|
||||
|
||||
// AppliedZoneCandidate is an account-level custom DNS zone reduced to the
|
||||
// per-peer decision the components calc still makes: include the zone only when
|
||||
// the peer belongs to one of its distribution groups. Record conversion is done
|
||||
// once at build time.
|
||||
type AppliedZoneCandidate struct {
|
||||
DistributionGroups []string
|
||||
Zone nmdata.CustomZone
|
||||
}
|
||||
|
||||
// PrivateServiceCandidate is a single private service's synthesized records,
|
||||
// carried per apex zone. The builder resolves proxy-cluster connectivity and
|
||||
// domain-suffix matching once; the calc merges the candidates whose AccessGroups
|
||||
// the peer belongs to, grouped by Zone.Domain.
|
||||
type PrivateServiceCandidate struct {
|
||||
AccessGroups []string
|
||||
Zone nmdata.CustomZone
|
||||
}
|
||||
18
shared/management/networkmap/nmdata/account_settings.go
Normal file
18
shared/management/networkmap/nmdata/account_settings.go
Normal file
@@ -0,0 +1,18 @@
|
||||
package nmdata
|
||||
|
||||
import "time"
|
||||
|
||||
// AccountSettingsInfo is the slim twin of types.AccountSettingsInfo.
|
||||
type AccountSettingsInfo struct {
|
||||
PeerLoginExpirationEnabled bool
|
||||
PeerLoginExpiration time.Duration
|
||||
PeerInactivityExpirationEnabled bool
|
||||
PeerInactivityExpiration time.Duration
|
||||
DNSDomain string
|
||||
IPv6EnabledGroups []string
|
||||
RoutingPeerDNSResolutionEnabled bool
|
||||
LazyConnectionEnabled bool
|
||||
AutoUpdateVersion string
|
||||
AutoUpdateAlways bool
|
||||
MetricsPushEnabled bool
|
||||
}
|
||||
18
shared/management/networkmap/nmdata/dns.go
Normal file
18
shared/management/networkmap/nmdata/dns.go
Normal file
@@ -0,0 +1,18 @@
|
||||
package nmdata
|
||||
|
||||
// SimpleRecord is the slim twin of dns.SimpleRecord.
|
||||
type SimpleRecord struct {
|
||||
Name string
|
||||
Type int
|
||||
Class string
|
||||
TTL int
|
||||
RData string
|
||||
}
|
||||
|
||||
// CustomZone is the slim twin of dns.CustomZone.
|
||||
type CustomZone struct {
|
||||
Domain string
|
||||
Records []SimpleRecord
|
||||
SearchDomainDisabled bool
|
||||
NonAuthoritative bool
|
||||
}
|
||||
6
shared/management/networkmap/nmdata/dns_settings.go
Normal file
6
shared/management/networkmap/nmdata/dns_settings.go
Normal file
@@ -0,0 +1,6 @@
|
||||
package nmdata
|
||||
|
||||
// DNSSettings is the slim twin of types.DNSSettings.
|
||||
type DNSSettings struct {
|
||||
DisabledManagementGroups []string
|
||||
}
|
||||
27
shared/management/networkmap/nmdata/group.go
Normal file
27
shared/management/networkmap/nmdata/group.go
Normal file
@@ -0,0 +1,27 @@
|
||||
package nmdata
|
||||
|
||||
import "slices"
|
||||
|
||||
const groupAllName = "All"
|
||||
|
||||
// Group is the slim twin of types.Group.
|
||||
type Group struct {
|
||||
ID string
|
||||
Name string
|
||||
PublicID string
|
||||
Peers []string
|
||||
Resources []Resource
|
||||
}
|
||||
|
||||
func (g *Group) IsGroupAll() bool {
|
||||
return g.Name == groupAllName
|
||||
}
|
||||
|
||||
func (g *Group) Copy() *Group {
|
||||
return &Group{
|
||||
ID: g.ID,
|
||||
Name: g.Name,
|
||||
PublicID: g.PublicID,
|
||||
Peers: slices.Clone(g.Peers),
|
||||
}
|
||||
}
|
||||
24
shared/management/networkmap/nmdata/nameserver.go
Normal file
24
shared/management/networkmap/nmdata/nameserver.go
Normal file
@@ -0,0 +1,24 @@
|
||||
package nmdata
|
||||
|
||||
import "net/netip"
|
||||
|
||||
// NameServerGroup is the slim twin of dns.NameServerGroup.
|
||||
type NameServerGroup struct {
|
||||
ID string
|
||||
PublicID string
|
||||
Name string
|
||||
Description string
|
||||
NameServers []NameServer
|
||||
Groups []string
|
||||
Primary bool
|
||||
Domains []string
|
||||
Enabled bool
|
||||
SearchDomainsEnabled bool
|
||||
}
|
||||
|
||||
// NameServer is the slim twin of dns.NameServer.
|
||||
type NameServer struct {
|
||||
IP netip.Addr
|
||||
NSType int
|
||||
Port int
|
||||
}
|
||||
16
shared/management/networkmap/nmdata/network.go
Normal file
16
shared/management/networkmap/nmdata/network.go
Normal file
@@ -0,0 +1,16 @@
|
||||
package nmdata
|
||||
|
||||
import "net"
|
||||
|
||||
// Network is the slim twin of types.Network.
|
||||
type Network struct {
|
||||
Identifier string
|
||||
Net net.IPNet
|
||||
NetV6 net.IPNet
|
||||
Dns string
|
||||
Serial int64
|
||||
}
|
||||
|
||||
func (n *Network) CurrentSerial() uint64 {
|
||||
return uint64(n.Serial)
|
||||
}
|
||||
18
shared/management/networkmap/nmdata/network_resource.go
Normal file
18
shared/management/networkmap/nmdata/network_resource.go
Normal file
@@ -0,0 +1,18 @@
|
||||
package nmdata
|
||||
|
||||
import "net/netip"
|
||||
|
||||
// NetworkResource is the slim twin of resources/types.NetworkResource.
|
||||
type NetworkResource struct {
|
||||
ID string
|
||||
NetworkID string
|
||||
AccountID string
|
||||
PublicID string
|
||||
Name string
|
||||
Description string
|
||||
Type string
|
||||
Address string // TODO: isn't persisted in the DB
|
||||
Domain string
|
||||
Prefix netip.Prefix
|
||||
Enabled bool
|
||||
}
|
||||
10
shared/management/networkmap/nmdata/network_router.go
Normal file
10
shared/management/networkmap/nmdata/network_router.go
Normal file
@@ -0,0 +1,10 @@
|
||||
package nmdata
|
||||
|
||||
// NetworkRouter is the slim twin of routers/types.NetworkRouter.
|
||||
type NetworkRouter struct {
|
||||
PublicID string
|
||||
PeerGroups []string
|
||||
Masquerade bool
|
||||
Metric int
|
||||
Enabled bool
|
||||
}
|
||||
126
shared/management/networkmap/nmdata/peer.go
Normal file
126
shared/management/networkmap/nmdata/peer.go
Normal file
@@ -0,0 +1,126 @@
|
||||
package nmdata
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
peerCapabilitySourcePrefixes int32 = 1
|
||||
peerCapabilityIPv6Overlay int32 = 2
|
||||
)
|
||||
|
||||
// Peer is the slim twin of peer.Peer.
|
||||
type Peer struct {
|
||||
ID string
|
||||
Key string
|
||||
SSHKey string
|
||||
DNSLabel string
|
||||
UserID string
|
||||
SSHEnabled bool
|
||||
LoginExpirationEnabled bool
|
||||
LastLogin *time.Time
|
||||
IP netip.Addr
|
||||
IPv6 netip.Addr
|
||||
RequiresApproval bool
|
||||
ExtraDNSLabels []string
|
||||
Meta PeerSystemMeta
|
||||
ProxyMeta ProxyMeta
|
||||
Location PeerLocation
|
||||
}
|
||||
|
||||
// ProxyMeta is the slim twin of peer.ProxyMeta.
|
||||
type ProxyMeta struct {
|
||||
Embedded bool
|
||||
}
|
||||
|
||||
// PeerSystemMeta is the slim twin of peer.PeerSystemMeta.
|
||||
type PeerSystemMeta struct {
|
||||
WtVersion string
|
||||
GoOS string
|
||||
OSVersion string
|
||||
KernelVersion string
|
||||
NetworkAddresses []NetworkAddress
|
||||
Files []File
|
||||
Capabilities []int32
|
||||
Flags Flags
|
||||
SyncMessageVersion int
|
||||
}
|
||||
|
||||
// Flags is the slim twin of peer.Flags.
|
||||
type Flags struct {
|
||||
ServerSSHAllowed bool
|
||||
DisableIPv6 bool
|
||||
}
|
||||
|
||||
// NetworkAddress is the slim twin of peer.NetworkAddress.
|
||||
type NetworkAddress struct {
|
||||
NetIP netip.Prefix
|
||||
}
|
||||
|
||||
// File is the slim twin of peer.File.
|
||||
type File struct {
|
||||
Path string
|
||||
ProcessIsRunning bool
|
||||
}
|
||||
|
||||
// PeerLocation is the slim twin of peer.Location.
|
||||
type PeerLocation struct {
|
||||
CountryCode string
|
||||
CityName string
|
||||
ConnectionIP net.IP
|
||||
}
|
||||
|
||||
func (p *Peer) HasCapability(capability int32) bool {
|
||||
return slices.Contains(p.Meta.Capabilities, capability)
|
||||
}
|
||||
|
||||
func (p *Peer) SupportsIPv6() bool {
|
||||
return !p.Meta.Flags.DisableIPv6 && p.HasCapability(peerCapabilityIPv6Overlay)
|
||||
}
|
||||
|
||||
func (p *Peer) SupportsSourcePrefixes() bool {
|
||||
return p.HasCapability(peerCapabilitySourcePrefixes)
|
||||
}
|
||||
|
||||
func (p *Peer) AddedWithSSOLogin() bool {
|
||||
return p.UserID != ""
|
||||
}
|
||||
|
||||
func (p *Peer) FQDN(dnsDomain string) string {
|
||||
if dnsDomain == "" {
|
||||
return ""
|
||||
}
|
||||
return p.DNSLabel + "." + dnsDomain
|
||||
}
|
||||
|
||||
func (p *Peer) GetLastLogin() time.Time {
|
||||
if p.LastLogin != nil {
|
||||
return *p.LastLogin
|
||||
}
|
||||
return time.Time{}
|
||||
}
|
||||
|
||||
// SessionExpiresAt mirrors peer.Peer.SessionExpiresAt.
|
||||
func (p *Peer) SessionExpiresAt(accountExpirationEnabled bool, expiresIn time.Duration) time.Time {
|
||||
if !accountExpirationEnabled || !p.AddedWithSSOLogin() || !p.LoginExpirationEnabled {
|
||||
return time.Time{}
|
||||
}
|
||||
last := p.GetLastLogin()
|
||||
if last.IsZero() {
|
||||
return time.Time{}
|
||||
}
|
||||
return last.Add(expiresIn).UTC()
|
||||
}
|
||||
|
||||
func (p *Peer) LoginExpired(expiresIn time.Duration) (bool, time.Duration) {
|
||||
if !p.AddedWithSSOLogin() || !p.LoginExpirationEnabled {
|
||||
return false, 0
|
||||
}
|
||||
expiresAt := p.GetLastLogin().Add(expiresIn)
|
||||
now := time.Now()
|
||||
timeLeft := expiresAt.Sub(now)
|
||||
return timeLeft <= 0, timeLeft
|
||||
}
|
||||
93
shared/management/networkmap/nmdata/policy.go
Normal file
93
shared/management/networkmap/nmdata/policy.go
Normal file
@@ -0,0 +1,93 @@
|
||||
package nmdata
|
||||
|
||||
const (
|
||||
policyRuleProtocolALL = "all"
|
||||
policyRuleProtocolTCP = "tcp"
|
||||
|
||||
defaultSSHPortString = "22"
|
||||
nativeSSHPortString = "22022"
|
||||
defaultSSHPortNumber uint16 = 22
|
||||
nativeSSHPortNumber uint16 = 22022
|
||||
)
|
||||
|
||||
// Policy is the slim twin of types.Policy.
|
||||
type Policy struct {
|
||||
ID string
|
||||
PublicID string
|
||||
Enabled bool
|
||||
SourcePostureChecks []string
|
||||
Rules []*PolicyRule
|
||||
}
|
||||
|
||||
// PolicyRule is the slim twin of types.PolicyRule.
|
||||
type PolicyRule struct {
|
||||
ID string
|
||||
PolicyID string
|
||||
Enabled bool
|
||||
Action string
|
||||
Protocol string
|
||||
Bidirectional bool
|
||||
Sources []string
|
||||
Destinations []string
|
||||
SourceResource Resource
|
||||
DestinationResource Resource
|
||||
Ports []string
|
||||
PortRanges []RulePortRange
|
||||
AuthorizedGroups map[string][]string
|
||||
AuthorizedUser string
|
||||
}
|
||||
|
||||
// RulePortRange is the slim twin of types.RulePortRange.
|
||||
type RulePortRange struct {
|
||||
Start uint16
|
||||
End uint16
|
||||
}
|
||||
|
||||
// Resource is the slim twin of types.Resource.
|
||||
type Resource struct {
|
||||
ID string
|
||||
Type string
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
// PolicyRuleImpliesLegacySSH is the twin-typed sibling of types.PolicyRuleImpliesLegacySSH.
|
||||
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
|
||||
}
|
||||
59
shared/management/networkmap/nmdata/posture.go
Normal file
59
shared/management/networkmap/nmdata/posture.go
Normal file
@@ -0,0 +1,59 @@
|
||||
package nmdata
|
||||
|
||||
const (
|
||||
checkActionAllow = "allow"
|
||||
checkActionDeny = "deny"
|
||||
)
|
||||
|
||||
// PostureChecks is the slim twin of posture.Checks.
|
||||
type PostureChecks struct {
|
||||
ID string
|
||||
Checks ChecksDefinition
|
||||
}
|
||||
|
||||
// ChecksDefinition is the slim twin of posture.ChecksDefinition.
|
||||
type ChecksDefinition struct {
|
||||
NBVersionCheck *NBVersionCheck
|
||||
OSVersionCheck *OSVersionCheck
|
||||
GeoLocationCheck *GeoLocationCheck
|
||||
PeerNetworkRangeCheck *PeerNetworkRangeCheck
|
||||
ProcessCheck *ProcessCheck
|
||||
}
|
||||
|
||||
type postureCheck interface {
|
||||
check(peer *Peer) (bool, error)
|
||||
}
|
||||
|
||||
// Passes reports whether the peer satisfies every check in this bundle. It
|
||||
// mirrors the server posture path: a check returning (false, _) — including on
|
||||
// an evaluation error — fails the bundle.
|
||||
func (pc *PostureChecks) Passes(peer *Peer) bool {
|
||||
for _, c := range pc.GetChecks() {
|
||||
valid, _ := c.check(peer)
|
||||
if !valid {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// GetChecks returns the initialized checks in the same order as posture.Checks.GetChecks.
|
||||
func (pc *PostureChecks) GetChecks() []postureCheck {
|
||||
var checks []postureCheck
|
||||
if pc.Checks.NBVersionCheck != nil {
|
||||
checks = append(checks, pc.Checks.NBVersionCheck)
|
||||
}
|
||||
if pc.Checks.OSVersionCheck != nil {
|
||||
checks = append(checks, pc.Checks.OSVersionCheck)
|
||||
}
|
||||
if pc.Checks.GeoLocationCheck != nil {
|
||||
checks = append(checks, pc.Checks.GeoLocationCheck)
|
||||
}
|
||||
if pc.Checks.PeerNetworkRangeCheck != nil {
|
||||
checks = append(checks, pc.Checks.PeerNetworkRangeCheck)
|
||||
}
|
||||
if pc.Checks.ProcessCheck != nil {
|
||||
checks = append(checks, pc.Checks.ProcessCheck)
|
||||
}
|
||||
return checks
|
||||
}
|
||||
45
shared/management/networkmap/nmdata/posture_geo_location.go
Normal file
45
shared/management/networkmap/nmdata/posture_geo_location.go
Normal file
@@ -0,0 +1,45 @@
|
||||
package nmdata
|
||||
|
||||
import "fmt"
|
||||
|
||||
// GeoLocation is the slim twin of posture.Location.
|
||||
type GeoLocation struct {
|
||||
CountryCode string
|
||||
CityName string
|
||||
}
|
||||
|
||||
// GeoLocationCheck is the slim twin of posture.GeoLocationCheck.
|
||||
type GeoLocationCheck struct {
|
||||
Locations []GeoLocation
|
||||
Action string
|
||||
}
|
||||
|
||||
func (g *GeoLocationCheck) check(peer *Peer) (bool, error) {
|
||||
if peer.Location.CountryCode == "" && peer.Location.CityName == "" {
|
||||
return false, fmt.Errorf("peer's location is not set")
|
||||
}
|
||||
|
||||
for _, loc := range g.Locations {
|
||||
if loc.CountryCode == peer.Location.CountryCode {
|
||||
if loc.CityName == "" || loc.CityName == peer.Location.CityName {
|
||||
switch g.Action {
|
||||
case checkActionDeny:
|
||||
return false, nil
|
||||
case checkActionAllow:
|
||||
return true, nil
|
||||
default:
|
||||
return false, fmt.Errorf("invalid geo location action: %s", g.Action)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if g.Action == checkActionDeny {
|
||||
return true, nil
|
||||
}
|
||||
if g.Action == checkActionAllow {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
return false, fmt.Errorf("invalid geo location action: %s", g.Action)
|
||||
}
|
||||
38
shared/management/networkmap/nmdata/posture_nb_version.go
Normal file
38
shared/management/networkmap/nmdata/posture_nb_version.go
Normal file
@@ -0,0 +1,38 @@
|
||||
package nmdata
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/hashicorp/go-version"
|
||||
)
|
||||
|
||||
// NBVersionCheck is the slim twin of posture.NBVersionCheck.
|
||||
type NBVersionCheck struct {
|
||||
MinVersion string
|
||||
}
|
||||
|
||||
func (n *NBVersionCheck) check(peer *Peer) (bool, error) {
|
||||
return meetsMinVersion(n.MinVersion, peer.Meta.WtVersion)
|
||||
}
|
||||
|
||||
func meetsMinVersion(minVer, peerVer string) (bool, error) {
|
||||
peerVer = sanitizeVersion(peerVer)
|
||||
minVer = sanitizeVersion(minVer)
|
||||
|
||||
peerNBVer, err := version.NewVersion(peerVer)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
constraints, err := version.NewConstraint(">= " + minVer)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
return constraints.Check(peerNBVer), nil
|
||||
}
|
||||
|
||||
func sanitizeVersion(v string) string {
|
||||
parts := strings.Split(v, "-")
|
||||
return parts[0]
|
||||
}
|
||||
62
shared/management/networkmap/nmdata/posture_network.go
Normal file
62
shared/management/networkmap/nmdata/posture_network.go
Normal file
@@ -0,0 +1,62 @@
|
||||
package nmdata
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/netip"
|
||||
)
|
||||
|
||||
// PeerNetworkRangeCheck is the slim twin of posture.PeerNetworkRangeCheck.
|
||||
type PeerNetworkRangeCheck struct {
|
||||
Action string
|
||||
Ranges []netip.Prefix
|
||||
}
|
||||
|
||||
func (p *PeerNetworkRangeCheck) check(peer *Peer) (bool, error) {
|
||||
peerPrefixes := make([]netip.Prefix, 0, len(peer.Meta.NetworkAddresses)+1)
|
||||
for _, peerNetAddr := range peer.Meta.NetworkAddresses {
|
||||
peerPrefixes = append(peerPrefixes, peerNetAddr.NetIP)
|
||||
}
|
||||
if connIP := peer.Location.ConnectionIP; len(connIP) > 0 {
|
||||
if addr, ok := netip.AddrFromSlice(connIP); ok {
|
||||
addr = addr.Unmap()
|
||||
peerPrefixes = append(peerPrefixes, netip.PrefixFrom(addr, addr.BitLen()))
|
||||
}
|
||||
}
|
||||
|
||||
if len(peerPrefixes) == 0 {
|
||||
return false, fmt.Errorf("peer's does not contain peer network range addresses")
|
||||
}
|
||||
|
||||
for _, peerPrefix := range peerPrefixes {
|
||||
for _, rangePrefix := range p.Ranges {
|
||||
if !prefixContains(rangePrefix, peerPrefix) {
|
||||
continue
|
||||
}
|
||||
switch p.Action {
|
||||
case checkActionDeny:
|
||||
return false, nil
|
||||
case checkActionAllow:
|
||||
return true, nil
|
||||
default:
|
||||
return false, fmt.Errorf("invalid peer network range check action: %s", p.Action)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if p.Action == checkActionDeny {
|
||||
return true, nil
|
||||
}
|
||||
if p.Action == checkActionAllow {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
return false, fmt.Errorf("invalid peer network range check action: %s", p.Action)
|
||||
}
|
||||
|
||||
func prefixContains(outer, inner netip.Prefix) bool {
|
||||
outer = outer.Masked()
|
||||
inner = inner.Masked()
|
||||
return outer.Bits() <= inner.Bits() &&
|
||||
outer.Addr().BitLen() == inner.Addr().BitLen() &&
|
||||
outer.Contains(inner.Addr())
|
||||
}
|
||||
79
shared/management/networkmap/nmdata/posture_os_version.go
Normal file
79
shared/management/networkmap/nmdata/posture_os_version.go
Normal file
@@ -0,0 +1,79 @@
|
||||
package nmdata
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/hashicorp/go-version"
|
||||
)
|
||||
|
||||
// MinVersionCheck is the slim twin of posture.MinVersionCheck.
|
||||
type MinVersionCheck struct {
|
||||
MinVersion string
|
||||
}
|
||||
|
||||
// MinKernelVersionCheck is the slim twin of posture.MinKernelVersionCheck.
|
||||
type MinKernelVersionCheck struct {
|
||||
MinKernelVersion string
|
||||
}
|
||||
|
||||
// OSVersionCheck is the slim twin of posture.OSVersionCheck.
|
||||
type OSVersionCheck struct {
|
||||
Android *MinVersionCheck
|
||||
Darwin *MinVersionCheck
|
||||
Ios *MinVersionCheck
|
||||
Linux *MinKernelVersionCheck
|
||||
Windows *MinKernelVersionCheck
|
||||
}
|
||||
|
||||
func (c *OSVersionCheck) check(peer *Peer) (bool, error) {
|
||||
switch peer.Meta.GoOS {
|
||||
case "android":
|
||||
return checkMinVersion(peer.Meta.OSVersion, c.Android)
|
||||
case "darwin":
|
||||
return checkMinVersion(peer.Meta.OSVersion, c.Darwin)
|
||||
case "ios":
|
||||
return checkMinVersion(peer.Meta.OSVersion, c.Ios)
|
||||
case "linux":
|
||||
kernelVersion := strings.Split(peer.Meta.KernelVersion, "-")[0]
|
||||
return checkMinKernelVersion(kernelVersion, c.Linux)
|
||||
case "windows":
|
||||
return checkMinKernelVersion(peer.Meta.KernelVersion, c.Windows)
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func checkMinVersion(peerVersion string, check *MinVersionCheck) (bool, error) {
|
||||
if check == nil {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
peerNBVersion, err := version.NewVersion(peerVersion)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
constraints, err := version.NewConstraint(">= " + check.MinVersion)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
return constraints.Check(peerNBVersion), nil
|
||||
}
|
||||
|
||||
func checkMinKernelVersion(peerVersion string, check *MinKernelVersionCheck) (bool, error) {
|
||||
if check == nil {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
peerNBVersion, err := version.NewVersion(peerVersion)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
constraints, err := version.NewConstraint(">= " + check.MinKernelVersion)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
return constraints.Check(peerNBVersion), nil
|
||||
}
|
||||
56
shared/management/networkmap/nmdata/posture_process.go
Normal file
56
shared/management/networkmap/nmdata/posture_process.go
Normal file
@@ -0,0 +1,56 @@
|
||||
package nmdata
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"slices"
|
||||
)
|
||||
|
||||
// Process is the slim twin of posture.Process.
|
||||
type Process struct {
|
||||
LinuxPath string
|
||||
MacPath string
|
||||
WindowsPath string
|
||||
}
|
||||
|
||||
// ProcessCheck is the slim twin of posture.ProcessCheck.
|
||||
type ProcessCheck struct {
|
||||
Processes []Process
|
||||
}
|
||||
|
||||
func (p *ProcessCheck) check(peer *Peer) (bool, error) {
|
||||
peerActiveProcesses := extractPeerActiveProcesses(peer.Meta.Files)
|
||||
|
||||
var pathSelector func(Process) string
|
||||
switch peer.Meta.GoOS {
|
||||
case "linux":
|
||||
pathSelector = func(process Process) string { return process.LinuxPath }
|
||||
case "darwin":
|
||||
pathSelector = func(process Process) string { return process.MacPath }
|
||||
case "windows":
|
||||
pathSelector = func(process Process) string { return process.WindowsPath }
|
||||
default:
|
||||
return false, fmt.Errorf("unsupported peer's operating system: %s", peer.Meta.GoOS)
|
||||
}
|
||||
|
||||
return p.areAllProcessesRunning(peerActiveProcesses, pathSelector), nil
|
||||
}
|
||||
|
||||
func (p *ProcessCheck) areAllProcessesRunning(activeProcesses []string, pathSelector func(Process) string) bool {
|
||||
for _, process := range p.Processes {
|
||||
path := pathSelector(process)
|
||||
if path == "" || !slices.Contains(activeProcesses, path) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func extractPeerActiveProcesses(files []File) []string {
|
||||
activeProcesses := make([]string, 0, len(files))
|
||||
for _, file := range files {
|
||||
if file.ProcessIsRunning {
|
||||
activeProcesses = append(activeProcesses, file.Path)
|
||||
}
|
||||
}
|
||||
return activeProcesses
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user