Compare commits

..

3 Commits

Author SHA1 Message Date
bcmmbaga
6725b02cbb Merge branch 'main' into add-atomic-cache-ops 2026-07-22 20:20:14 +03:00
bcmmbaga
da19dcf480 prevent concurrent JWT reuse 2026-07-22 20:08:48 +03:00
bcmmbaga
6426d6f03f add atomic SetNX cache operation
Split the memory and Redis cache implementations into separate files and
provide backend-native atomic set-if-absent support.
2026-07-22 20:03:03 +03:00
136 changed files with 2575 additions and 9876 deletions

View File

@@ -51,9 +51,6 @@ jobs:
# token (and URL, for gateways) is unset, so partial coverage is fine.
OPENAI_TOKEN: ${{ secrets.E2E_OPENAI_TOKEN }}
ANTHROPIC_TOKEN: ${{ secrets.E2E_ANTHROPIC_TOKEN }}
# Moonshot AI platform key (platform.kimi.ai); drives both Kimi wire
# shapes (OpenAI /v1 and Anthropic /anthropic) through kimi_api.
KIMI_TOKEN: ${{ secrets.E2E_KIMI_TOKEN }}
VERCEL_URL: ${{ secrets.E2E_VERCEL_URL }}
VERCEL_TOKEN: ${{ secrets.E2E_VERCEL_TOKEN }}
OPENROUTER_URL: ${{ secrets.E2E_OPENROUTER_URL }}

View File

@@ -569,17 +569,17 @@ func Test_ConnectPeers(t *testing.T) {
if err != nil {
t.Fatal(err)
}
// The peers use userspace WireGuard (stdnet transport). A tight busy-loop
// here starves the wireguard-go goroutines that process the handshake, so
// poll on a ticker instead and yield the CPU between checks. WireGuard also
// only retries a lost handshake initiation every REKEY_TIMEOUT (5s), which
// is why the overall wait can occasionally stretch to tens of seconds.
// todo: investigate why in some tests execution we need 30s
timeout := 30 * time.Second
timeoutChannel := time.After(timeout)
ticker := time.NewTicker(500 * time.Millisecond)
defer ticker.Stop()
for {
select {
case <-timeoutChannel:
t.Fatalf("waiting for peer handshake timeout after %s", timeout.String())
default:
}
peer, gpErr := getPeer(peer1ifaceName, peer2Key.PublicKey().String())
if gpErr != nil {
t.Fatal(gpErr)
@@ -588,12 +588,6 @@ func Test_ConnectPeers(t *testing.T) {
t.Log("peers successfully handshake")
break
}
select {
case <-timeoutChannel:
t.Fatalf("waiting for peer handshake timeout after %s", timeout.String())
case <-ticker.C:
}
}
}

View File

@@ -49,21 +49,11 @@ type ConnMgr struct {
// engine.syncMsgMux; all other reads stay under engine.syncMsgMux only.
lazyConnMgrMu sync.RWMutex
// reconcileRoutedIPs re-applies a peer's routed allowed IPs after its lazy wake endpoint is
// (re)armed (Mode A at arm time). Injected by the engine; nil disables the reconcile.
reconcileRoutedIPs func(peerKey string) error
wg sync.WaitGroup
lazyCtx context.Context
lazyCtxCancel context.CancelFunc
}
// SetRoutedIPsReconciler injects the callback used to re-apply a peer's routed allowed IPs when
// its lazy wake endpoint is (re)armed. Must be called before the lazy manager starts.
func (e *ConnMgr) SetRoutedIPsReconciler(fn func(peerKey string) error) {
e.reconcileRoutedIPs = fn
}
func NewConnMgr(engineConfig *EngineConfig, statusRecorder *peer.Status, peerStore *peerstore.Store, iface lazyconn.WGIface) *ConnMgr {
e := &ConnMgr{
peerStore: peerStore,
@@ -301,7 +291,6 @@ func (e *ConnMgr) Close() {
func (e *ConnMgr) initLazyManager(engineCtx context.Context) {
cfg := manager.Config{
InactivityThreshold: inactivityThresholdEnv(),
ReconcileAllowedIPs: e.reconcileRoutedIPs,
}
e.lazyConnMgrMu.Lock()

View File

@@ -52,14 +52,11 @@ int xdp_dns_fwd(struct iphdr *ip, struct udphdr *udp) {
if (udp->dest == GENERAL_DNS_PORT && ip->daddr == dns_ip) {
udp->dest = dns_port;
// Clear the now-stale checksum; zero means "not computed" for IPv4.
udp->check = 0;
return XDP_PASS;
}
if (udp->source == dns_port && ip->saddr == dns_ip) {
udp->source = GENERAL_DNS_PORT;
udp->check = 0;
return XDP_PASS;
}

View File

@@ -50,11 +50,5 @@ int xdp_wg_proxy(struct iphdr *ip, struct udphdr *udp) {
__be16 new_dst_port = htons(proxy_port);
udp->dest = new_dst_port;
udp->source = new_src_port;
// The ports are covered by the UDP checksum. This is an IPv4 loopback hop
// and the payload is already integrity-protected, so clear the checksum (a
// zero UDP checksum means "not computed" for IPv4) rather than leave a
// stale value the kernel would drop as UDP_CSUM.
udp->check = 0;
return XDP_PASS;
}

View File

@@ -663,12 +663,6 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL)
iceCfg := e.createICEConfig()
e.connMgr = NewConnMgr(e.config, e.statusRecorder, e.peerStore, wgIface)
e.connMgr.SetRoutedIPsReconciler(func(peerKey string) error {
if e.routeManager == nil {
return nil
}
return e.routeManager.ReconcilePeerAllowedIPs(peerKey)
})
e.connMgr.Start(e.ctx)
// Wire DNS-time lazy-connection warm-up now that the connection manager

View File

@@ -29,11 +29,6 @@ type managedPeer struct {
type Config struct {
InactivityThreshold *time.Duration
// ReconcileAllowedIPs re-applies a peer's routed allowed IPs after its wake endpoint is
// armed. The activity listener creates the wake peer with the overlay /32 only; without the
// routed prefixes WireGuard would not steer subnet-bound traffic to the wake endpoint, so an
// idle routing peer could never be woken by that traffic. Optional; nil disables the reconcile.
ReconcileAllowedIPs func(peerKey string) error
}
// Manager manages lazy connections
@@ -61,9 +56,6 @@ type Manager struct {
peerToHAGroups map[string][]route.HAUniqueID // peer ID -> HA groups they belong to
haGroupToPeers map[route.HAUniqueID][]string // HA group -> peer IDs in the group
routesMu sync.RWMutex
// reconcileAllowedIPs re-applies a peer's routed allowed IPs after its wake endpoint is armed.
reconcileAllowedIPs func(peerKey string) error
}
// NewManager creates a new lazy connection manager
@@ -81,7 +73,6 @@ func NewManager(config Config, engineCtx context.Context, peerStore *peerstore.S
activityManager: activity.NewManager(wgIface),
peerToHAGroups: make(map[string][]route.HAUniqueID),
haGroupToPeers: make(map[route.HAUniqueID][]string),
reconcileAllowedIPs: config.ReconcileAllowedIPs,
}
if wgIface.IsUserspaceBind() {
@@ -210,7 +201,7 @@ func (m *Manager) AddPeer(peerCfg lazyconn.PeerConfig) (bool, error) {
return false, nil
}
if err := m.armActivityListener(peerCfg); err != nil {
if err := m.activityManager.MonitorPeerActivity(peerCfg); err != nil {
return false, err
}
@@ -297,7 +288,7 @@ func (m *Manager) DeactivatePeer(peerID peerid.ConnID) {
m.inactivityManager.RemovePeer(mp.peerCfg.PublicKey)
if err := m.armActivityListener(*mp.peerCfg); err != nil {
if err := m.activityManager.MonitorPeerActivity(*mp.peerCfg); err != nil {
mp.peerCfg.Log.Errorf("failed to create activity monitor: %v", err)
return
}
@@ -474,31 +465,6 @@ func (m *Manager) close() {
}
// shouldDeferIdleForHA checks if peer should stay connected due to HA group requirements
// armRoutedAllowedIPs re-applies the peer's routed allowed IPs onto its freshly armed wake
// endpoint. The activity listener creates the wake peer with the overlay /32 only, so without
// this the routed prefixes would be missing and traffic to a routed subnet could not wake the
// idle routing peer. It is a no-op when no reconciler is configured.
// armActivityListener (re)arms the peer's wake endpoint via the activity manager and then
// re-applies its routed allowed IPs, so traffic to a routed subnet can wake an idle routing
// peer. The routed prefixes must be re-applied after the wake endpoint exists because the
// listener creates it with the overlay /32 only.
func (m *Manager) armActivityListener(peerCfg lazyconn.PeerConfig) error {
if err := m.activityManager.MonitorPeerActivity(peerCfg); err != nil {
return err
}
m.armRoutedAllowedIPs(&peerCfg)
return nil
}
func (m *Manager) armRoutedAllowedIPs(peerCfg *lazyconn.PeerConfig) {
if m.reconcileAllowedIPs == nil {
return
}
if err := m.reconcileAllowedIPs(peerCfg.PublicKey); err != nil {
peerCfg.Log.Errorf("failed to reconcile routed allowed IPs on wake endpoint: %v", err)
}
}
func (m *Manager) shouldDeferIdleForHA(inactivePeers map[string]struct{}, peerID string) bool {
m.routesMu.RLock()
defer m.routesMu.RUnlock()
@@ -611,7 +577,7 @@ func (m *Manager) onPeerInactivityTimedOut(peerIDs map[string]struct{}) {
mp.peerCfg.Log.Infof("start activity monitor")
if err := m.armActivityListener(*mp.peerCfg); err != nil {
if err := m.activityManager.MonitorPeerActivity(*mp.peerCfg); err != nil {
mp.peerCfg.Log.Errorf("failed to create activity monitor: %v", err)
continue
}

View File

@@ -61,7 +61,6 @@ type Manager interface {
InitialRouteRange() []string
SetFirewall(firewall.Manager) error
SetDNSForwarderPort(port uint16)
ReconcilePeerAllowedIPs(peerKey string) error
Stop(stateManager *statemanager.Manager)
}
@@ -233,30 +232,6 @@ func (m *DefaultManager) setupRefCounters(useNoop bool) {
)
}
// ReconcilePeerAllowedIPs re-applies every routed allowed IP currently tracked for the peer
// onto the WireGuard device. The allowed-IP refcounter only calls its AddFunc (which pushes to
// the device) on a prefix's 0->1 transition, so a peer whose device entry was rebuilt without a
// matching refcounter change — e.g. a lazy connection cycling through idle->wake, which recreates
// the WireGuard peer with the overlay /32 only — ends up missing routed prefixes the refcounter
// still considers installed, and nothing retries. Calling this when the peer's WireGuard entry is
// (re)created restores convergence. It is add-only and idempotent: AddAllowedIP is update-only, so
// prefixes are re-added to an existing peer and an absent peer is left untouched.
func (m *DefaultManager) ReconcilePeerAllowedIPs(peerKey string) error {
if m.allowedIPsRefCounter == nil {
return nil
}
return m.allowedIPsRefCounter.ReapplyMatching(
func(out string) bool { return out == peerKey },
func(prefix netip.Prefix) error {
if err := m.wgInterface.AddAllowedIP(peerKey, prefix); err != nil {
return fmt.Errorf("add allowed IP %s for peer %s: %w", prefix, peerKey, err)
}
return nil
},
)
}
// Init sets up the routing
func (m *DefaultManager) Init() error {
m.routeSelector = m.initSelector()

View File

@@ -112,11 +112,6 @@ func (m *MockManager) SetFirewall(firewall.Manager) error {
func (m *MockManager) SetDNSForwarderPort(port uint16) {
}
// ReconcilePeerAllowedIPs mock implementation of ReconcilePeerAllowedIPs from Manager interface
func (m *MockManager) ReconcilePeerAllowedIPs(peerKey string) error {
return nil
}
// Stop mock implementation of Stop from Manager interface
func (m *MockManager) Stop(stateManager *statemanager.Manager) {
if m.StopFunc != nil {

View File

@@ -1,90 +0,0 @@
//go:build !windows
package routemanager
import (
"net"
"net/netip"
"sync"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.zx2c4.com/wireguard/tun/netstack"
"github.com/netbirdio/netbird/client/iface/device"
"github.com/netbirdio/netbird/client/iface/wgaddr"
"github.com/netbirdio/netbird/client/internal/routemanager/refcounter"
)
// reconcileWGMock is a minimal iface.WGIface that only records AddAllowedIP calls; every other
// method is an inert stub because ReconcilePeerAllowedIPs exercises none of them.
type reconcileWGMock struct {
mu sync.Mutex
adds map[string][]netip.Prefix
}
func (m *reconcileWGMock) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
m.mu.Lock()
defer m.mu.Unlock()
if m.adds == nil {
m.adds = map[string][]netip.Prefix{}
}
m.adds[peerKey] = append(m.adds[peerKey], allowedIP)
return nil
}
func (m *reconcileWGMock) added(peerKey string) []netip.Prefix {
m.mu.Lock()
defer m.mu.Unlock()
return m.adds[peerKey]
}
func (m *reconcileWGMock) RemoveAllowedIP(string, netip.Prefix) error { return nil }
func (m *reconcileWGMock) Name() string { return "utun-test" }
func (m *reconcileWGMock) Address() wgaddr.Address { return wgaddr.Address{} }
func (m *reconcileWGMock) ToInterface() *net.Interface { return nil }
func (m *reconcileWGMock) IsUserspaceBind() bool { return false }
func (m *reconcileWGMock) GetFilter() device.PacketFilter { return nil }
func (m *reconcileWGMock) GetDevice() *device.FilteredDevice { return nil }
func (m *reconcileWGMock) GetNet() *netstack.Net { return nil }
// TestReconcilePeerAllowedIPs verifies the declarative reconcile re-applies every routed prefix
// tracked for the peer (self-heal, independent of refcount level) and stays scoped to that peer.
func TestReconcilePeerAllowedIPs(t *testing.T) {
wg := &reconcileWGMock{}
m := &DefaultManager{wgInterface: wg}
m.allowedIPsRefCounter = refcounter.New[netip.Prefix, string, string](
func(_ netip.Prefix, peerKey string) (string, error) { return peerKey, nil },
func(netip.Prefix, string) error { return nil },
)
peerA1 := netip.MustParsePrefix("10.0.0.0/24")
peerA2 := netip.MustParsePrefix("10.1.0.0/24")
peerB1 := netip.MustParsePrefix("10.2.0.0/24")
for prefix, peer := range map[netip.Prefix]string{peerA1: "peerA", peerA2: "peerA", peerB1: "peerB"} {
_, err := m.allowedIPsRefCounter.Increment(prefix, peer)
require.NoError(t, err)
}
// Extra reference: reconcile must still re-apply the prefix even though its refcount never
// hit 0 again (the exact case the plain incremental path skips).
_, err := m.allowedIPsRefCounter.Increment(peerA1, "peerA")
require.NoError(t, err)
require.NoError(t, m.ReconcilePeerAllowedIPs("peerA"))
assert.ElementsMatch(t, []netip.Prefix{peerA1, peerA2}, wg.added("peerA"),
"reconcile must re-apply all routed prefixes of the peer")
assert.Empty(t, wg.added("peerB"), "reconcile must not touch another peer's prefixes")
}
// TestReconcilePeerAllowedIPsNoCounter verifies reconcile is a safe no-op before the refcounter is
// set up.
func TestReconcilePeerAllowedIPsNoCounter(t *testing.T) {
wg := &reconcileWGMock{}
m := &DefaultManager{wgInterface: wg}
require.NoError(t, m.ReconcilePeerAllowedIPs("peerA"))
assert.Empty(t, wg.added("peerA"))
}

View File

@@ -94,26 +94,6 @@ func (rm *Counter[Key, I, O]) Get(key Key) (Ref[O], bool) {
return ref, ok
}
// ReapplyMatching calls apply for every key whose stored Out satisfies pred, holding the
// counter lock for the whole pass. Running apply under the lock keeps it atomic with respect
// to Increment/Decrement: a prefix dropped to zero is removed from the map (and had its
// RemoveFunc called) before this pass observes it, so a stale key can never be re-applied.
// pred and apply are invoked under the lock, so they must not call back into the counter.
func (rm *Counter[Key, I, O]) ReapplyMatching(pred func(out O) bool, apply func(key Key) error) error {
rm.mu.Lock()
defer rm.mu.Unlock()
var merr *multierror.Error
for key, ref := range rm.refCountMap {
if pred(ref.Out) {
if err := apply(key); err != nil {
merr = multierror.Append(merr, err)
}
}
}
return nberrors.FormatErrorOrNil(merr)
}
// Increment increments the reference count for the given key.
// If this is the first reference to the key, the AddFunc is called.
func (rm *Counter[Key, I, O]) Increment(key Key, in I) (Ref[O], error) {

View File

@@ -1,47 +0,0 @@
package refcounter
import (
"net/netip"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// TestReapplyMatching verifies ReapplyMatching invokes apply for exactly the keys whose stored
// Out satisfies the predicate (no duplicates for multiply-referenced keys) — the primitive
// ReconcilePeerAllowedIPs relies on to re-apply a single peer's routed prefixes.
func TestReapplyMatching(t *testing.T) {
rc := New[netip.Prefix, string, string](
func(_ netip.Prefix, peerKey string) (string, error) { return peerKey, nil },
func(netip.Prefix, string) error { return nil },
)
peerA1 := netip.MustParsePrefix("10.0.0.0/24")
peerA2 := netip.MustParsePrefix("10.1.0.0/24")
peerB1 := netip.MustParsePrefix("10.2.0.0/24")
for prefix, peer := range map[netip.Prefix]string{peerA1: "peerA", peerA2: "peerA", peerB1: "peerB"} {
_, err := rc.Increment(prefix, peer)
require.NoError(t, err)
}
// a second reference must not make the key applied twice
_, err := rc.Increment(peerA1, "peerA")
require.NoError(t, err)
var applied []netip.Prefix
err = rc.ReapplyMatching(
func(out string) bool { return out == "peerA" },
func(key netip.Prefix) error { applied = append(applied, key); return nil },
)
require.NoError(t, err)
assert.ElementsMatch(t, []netip.Prefix{peerA1, peerA2}, applied)
var none []netip.Prefix
err = rc.ReapplyMatching(
func(out string) bool { return out == "missing" },
func(key netip.Prefix) error { none = append(none, key); return nil },
)
require.NoError(t, err)
assert.Empty(t, none)
}

View File

@@ -20,15 +20,14 @@ import (
// covers whatever credentials are present (source ~/.llm-keys locally / set the
// Actions secrets in CI).
type providerCase struct {
name string
catalogID string
upstream string
apiKey string
model string // body model (chat/messages) or path model@version (vertex)
kind string // harness.WireChat, harness.WireMessages, or harness.WireVertex
project string // vertex only: GCP project for the rawPredict path
region string // vertex only: GCP region for the rawPredict path
pathPrefix string // base-URL path prefix the agent carries (e.g. "/anthropic" for Kimi)
name string
catalogID string
upstream string
apiKey string
model string // body model (chat/messages) or path model@version (vertex)
kind string // harness.WireChat, harness.WireMessages, or harness.WireVertex
project string // vertex only: GCP project for the rawPredict path
region string // vertex only: GCP region for the rawPredict path
}
// availableProviders builds the matrix from the provider env vars that are set.
@@ -40,24 +39,6 @@ func availableProviders() []providerCase {
if k := os.Getenv("ANTHROPIC_TOKEN"); k != "" {
ps = append(ps, providerCase{name: "anthropic", catalogID: "anthropic_api", upstream: "https://api.anthropic.com", apiKey: k, model: "claude-haiku-4-5", kind: harness.WireMessages})
}
if k := os.Getenv("KIMI_TOKEN"); k != "" {
// Kimi (Moonshot AI) serves two body shapes from the same key: OpenAI
// Chat Completions on the bare host (/v1/...) and the Anthropic
// Messages API under the /anthropic path prefix (the endpoint
// Moonshot's Claude Code guide uses). The provider keeps the bare
// default upstream and the AGENT carries the /anthropic prefix in
// its base URL — exactly the documented Claude Code / Kimi CLI
// setup (ANTHROPIC_BASE_URL=https://<endpoint>/anthropic) — so one
// provider serves both shapes and the prefix rides through to
// Moonshot. Run the Anthropic shape, the flagship Claude Code path;
// the OpenAI wire shape is covered live by the other chat-shaped
// matrix providers, and Kimi-over-chat passed with kimi-k3 before
// the single-model constraint surfaced (run #73 on the kimi feature
// branch). The platform serves this account exactly ONE model —
// kimi-k3 (kimi-k2-thinking and even kimi-latest return
// resource_not_found_error on both surfaces).
ps = append(ps, providerCase{name: "kimi", catalogID: "kimi_api", upstream: "https://api.moonshot.ai", apiKey: k, model: "kimi-k3", kind: harness.WireMessages, pathPrefix: "/anthropic"})
}
if k, u := os.Getenv("VERCEL_TOKEN"), os.Getenv("VERCEL_URL"); k != "" && u != "" {
ps = append(ps, providerCase{name: "vercel", catalogID: "vercel_ai_gateway", upstream: u, apiKey: k, model: "openai/gpt-4o-mini", kind: harness.WireChat})
}
@@ -103,18 +84,12 @@ func availableProviders() []providerCase {
}
}
// Bedrock: path-routed, bearer auth. Model is the FULL cross-region
// inference-profile id exactly as AWS issues it — region-family prefix
// plus the date/version suffix. A bare or wrong-region id makes Bedrock
// reject the request with "The provided model identifier is invalid"
// before any inference runs. The proxy normalizes this id to the catalog
// key (anthropic.claude-haiku-4-5) for routing/pricing/allowlists.
// Defaults pair eu-central-1 with the eu.* profile; AWS_REGION overrides
// the region and the prefix follows its family.
// Bedrock: path-routed, bearer auth. Model is a cross-region inference
// profile id (distinct string from the first-party Anthropic case).
if k := os.Getenv("AWS_BEARER_TOKEN_BEDROCK"); k != "" {
region := os.Getenv("AWS_REGION")
if region == "" {
region = "eu-central-1"
region = "us-east-1"
}
// A valid Bedrock inference-profile id (region prefix + date + version),
// overridable per account. `global.` profiles can be invoked from any
@@ -241,7 +216,8 @@ func TestProvidersMatrix(t *testing.T) {
t.Cleanup(func() { _ = cl.Terminate(context.Background()) })
require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management")
// Probe first: the GET resolves the endpoint (DNS error fails) and its first packet wakes the lazy proxy peer, so WaitProxyPeer sees it connected; any HTTP status counts.
// Resolve first: the DNS lookup triggers the lazy-connection warm-up, waking
// the proxy peer so WaitProxyPeer then observes it connected.
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
require.NoError(t, err, "resolve agent-network endpoint to proxy IP")
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {
@@ -271,7 +247,7 @@ func TestProvidersMatrix(t *testing.T) {
case harness.WireBedrock:
c, b, cerr = cl.Bedrock(ctx, settings.Endpoint, proxyIP, pc.model, "Reply with exactly: pong", sessionID)
default:
c, b, cerr = cl.ChatPrefixed(ctx, settings.Endpoint, proxyIP, pc.pathPrefix, pc.kind, pc.model, "Reply with exactly: pong", sessionID)
c, b, cerr = cl.Chat(ctx, settings.Endpoint, proxyIP, pc.kind, pc.model, "Reply with exactly: pong", sessionID)
}
if cerr == nil {
code, body = c, b

View File

@@ -52,9 +52,7 @@ func catalogModel(pc providerCase) string {
func disallowedModel(pc providerCase) string {
switch pc.kind {
case harness.WireBedrock:
// Same profile prefix as the allowed model so only the model name
// differs; the guardrail must deny it before it reaches AWS.
return strings.SplitN(pc.model, ".", 2)[0] + ".anthropic.claude-opus-4-8"
return "us.anthropic.claude-opus-4-8"
case harness.WireVertex:
return "claude-opus-4-8@20250101"
default:
@@ -74,7 +72,7 @@ func sendModel(ctx context.Context, t *testing.T, cl *harness.Client, endpoint,
case harness.WireVertex:
code, _, err = cl.Vertex(ctx, endpoint, proxyIP, pc.project, pc.region, model, "Reply with exactly: pong", "")
default:
code, _, err = cl.ChatPrefixed(ctx, endpoint, proxyIP, pc.pathPrefix, pc.kind, model, "Reply with exactly: pong", "")
code, _, err = cl.Chat(ctx, endpoint, proxyIP, pc.kind, model, "Reply with exactly: pong", "")
}
require.NoError(t, err, "request must reach the proxy for %s", pc.name)
return code
@@ -166,7 +164,8 @@ func TestModelAllowlistEnforced(t *testing.T) {
t.Cleanup(func() { _ = cl.Terminate(context.Background()) })
require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management")
// Probe first: the GET resolves the endpoint (DNS error fails) and its first packet wakes the lazy proxy peer, so WaitProxyPeer sees it connected; any HTTP status counts.
// Resolve first: the DNS lookup triggers the lazy-connection warm-up, waking
// the proxy peer so WaitProxyPeer then observes it connected.
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
require.NoError(t, err, "resolve agent-network endpoint to proxy IP")
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {

View File

@@ -104,7 +104,8 @@ func TestProviderSkipTLSVerification(t *testing.T) {
t.Cleanup(func() { _ = cl.Terminate(context.Background()) })
require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management")
// Probe first: the GET resolves the endpoint (DNS error fails) and its first packet wakes the lazy proxy peer, so WaitProxyPeer sees it connected; any HTTP status counts.
// Resolve first: the DNS lookup triggers the lazy-connection warm-up, waking
// the proxy peer so WaitProxyPeer then observes it connected.
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
require.NoError(t, err, "resolve endpoint to proxy IP")
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {

View File

@@ -106,7 +106,8 @@ func TestVLLMProvider(t *testing.T) {
t.Cleanup(func() { _ = cl.Terminate(context.Background()) })
require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management")
// Probe first: the GET resolves the endpoint (DNS error fails) and its first packet wakes the lazy proxy peer, so WaitProxyPeer sees it connected; any HTTP status counts.
// Resolve first: the DNS lookup triggers the lazy-connection warm-up, waking
// the proxy peer so WaitProxyPeer then observes it connected.
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
require.NoError(t, err, "resolve endpoint to proxy IP")
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {

View File

@@ -4,7 +4,6 @@ package harness
import (
"context"
"errors"
"fmt"
"io"
"os/exec"
@@ -168,54 +167,22 @@ func (cl *Client) pollStatus(ctx context.Context, timeout time.Duration, want st
return fmt.Errorf("timed out waiting for %q; last status:\n%s", want, last)
}
const (
// curlExitCouldNotResolve is curl's exit code for a DNS resolution failure, distinct from connection-level failures.
curlExitCouldNotResolve = 6
// dnsProbeRetryWindow bounds DNS-failure retries: the synthesized zone lands a beat after management connects, so early NXDOMAIN is propagation; a zone still absent after this window is a real failure.
dnsProbeRetryWindow = 30 * time.Second
dnsProbeRetryInterval = 2 * time.Second
)
// ResolveProxyIP GETs https://<endpoint>/ from the client's netns: any HTTP status proves DNS + tunnel and wakes the lazy proxy peer; only DNS failures retry, within dnsProbeRetryWindow. Returns the connected IP for --resolve pinning.
// ResolveProxyIP resolves the agent-network endpoint to the proxy peer's
// NetBird IP from inside the client (via magic DNS).
func (cl *Client) ResolveProxyIP(ctx context.Context, endpoint string) (string, error) {
args := []string{
"run", "--rm",
"--network", "container:" + cl.container.GetContainerID(),
curlImage,
"-ksS", "-o", "/dev/null",
"--connect-timeout", "30", "--max-time", "60",
"-w", "%{remote_ip}",
"https://" + endpoint + "/",
code, reader, err := cl.container.Exec(ctx, []string{"getent", "hosts", endpoint}, tcexec.Multiplexed())
if err != nil {
return "", err
}
deadline := time.Now().Add(dnsProbeRetryWindow)
for {
cmd := exec.CommandContext(ctx, "docker", args...)
var stdout, stderr strings.Builder
cmd.Stdout = &stdout
cmd.Stderr = &stderr
err := cmd.Run()
if err == nil {
ip := strings.TrimSpace(stdout.String())
if ip == "" {
return "", fmt.Errorf("got an HTTP response from %s but no remote IP", endpoint)
}
return ip, nil
}
var exitErr *exec.ExitError
if !errors.As(err, &exitErr) || exitErr.ExitCode() != curlExitCouldNotResolve {
return "", fmt.Errorf("no HTTP response from %s: %w (%s)", endpoint, err, strings.TrimSpace(stderr.String()))
}
dnsErr := fmt.Errorf("DNS resolution failed for %s: %s", endpoint, strings.TrimSpace(stderr.String()))
if time.Until(deadline) < dnsProbeRetryInterval {
return "", dnsErr
}
select {
case <-ctx.Done():
return "", fmt.Errorf("%w (%w)", dnsErr, ctx.Err())
case <-time.After(dnsProbeRetryInterval):
}
out, _ := io.ReadAll(reader)
if code != 0 {
return "", fmt.Errorf("getent hosts %s exited %d", endpoint, code)
}
fields := strings.Fields(string(out))
if len(fields) == 0 {
return "", fmt.Errorf("no address for %s", endpoint)
}
return fields[0], nil
}
// Wire shapes for Chat.
@@ -239,17 +206,6 @@ const (
// the wire shape: WireChat (OpenAI) or WireMessages (Anthropic). A non-empty
// sessionID is sent as the universal x-session-id header the proxy records.
func (cl *Client) Chat(ctx context.Context, endpoint, proxyIP, kind, model, prompt, sessionID string) (int, string, error) {
return cl.ChatPrefixed(ctx, endpoint, proxyIP, "", kind, model, prompt, sessionID)
}
// ChatPrefixed is Chat with a base-URL path prefix prepended to the wire
// path, mirroring agents whose base URL carries a shape-selecting prefix that
// rides through to the upstream — e.g. Claude Code against a Kimi provider
// sets ANTHROPIC_BASE_URL=https://<endpoint>/anthropic so the proxy forwards
// /anthropic/v1/messages to Moonshot's Anthropic surface while the provider's
// upstream URL stays the bare https://api.moonshot.ai. Empty prefix is plain
// Chat.
func (cl *Client) ChatPrefixed(ctx context.Context, endpoint, proxyIP, pathPrefix, kind, model, prompt, sessionID string) (int, string, error) {
var path, body string
var headers []string
switch kind {
@@ -261,7 +217,7 @@ func (cl *Client) ChatPrefixed(ctx context.Context, endpoint, proxyIP, pathPrefi
path = "/v1/chat/completions"
body = fmt.Sprintf(`{"model":%q,"messages":[{"role":"user","content":%q}]}`, model, prompt)
}
return cl.post(ctx, endpoint, proxyIP, pathPrefix+path, body, withSessionID(headers, sessionID))
return cl.post(ctx, endpoint, proxyIP, path, body, withSessionID(headers, sessionID))
}
// Vertex issues an Anthropic-on-Vertex rawPredict POST over the tunnel. Unlike

2
go.mod
View File

@@ -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-20260803100840-78e79ba20f87
github.com/netbirdio/management-integrations/integrations v0.0.0-20260416123949-2355d972be42
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
View File

@@ -486,8 +486,6 @@ 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=

View File

@@ -234,6 +234,9 @@ init_environment() {
NETBIRD_LICENSE_KEY=$(read_secret "Enter license key (input hidden)")
GHCR_USERNAME="netbirdExtAccess1"
GHCR_TOKEN=$(read_secret "Enter GHCR token (input hidden)")
POSTGRES_USER="netbird"
POSTGRES_DB="netbird"
POSTGRES_PASSWORD=$(rand_secret)
@@ -260,6 +263,10 @@ init_environment() {
install -m 600 /dev/null config.yaml
render_config_yaml >> config.yaml
echo "Logging in to ghcr.io ..."
printf '%s' "$GHCR_TOKEN" | docker login ghcr.io -u "$GHCR_USERNAME" --password-stdin
unset GHCR_TOKEN
echo ""
echo "Pulling images ..."
$DOCKER_COMPOSE_COMMAND pull

View File

@@ -490,6 +490,8 @@ init_migration() {
echo ""
echo "Step 1: Image swap (community → Enterprise). License key required."
NB_LICENSE_KEY=$(read_secret " License key")
GHCR_USERNAME="netbirdExtAccess1"
GHCR_TOKEN=$(read_secret " GHCR token (input hidden)")
# Step 2 — optional
echo ""
@@ -586,6 +588,11 @@ apply_changes() {
fi
} >> "$ENV_FILE"
echo ""
echo "Logging in to ghcr.io ..."
printf '%s' "$GHCR_TOKEN" | docker login ghcr.io -u "$GHCR_USERNAME" --password-stdin
unset GHCR_TOKEN
echo ""
echo "Pulling enterprise images ..."
$DOCKER_COMPOSE_COMMAND pull

View File

@@ -1,39 +0,0 @@
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"}}},
)
}

View File

@@ -1,115 +0,0 @@
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}`)
}

View File

@@ -18,7 +18,6 @@ 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"
@@ -31,8 +30,6 @@ 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"
@@ -64,8 +61,6 @@ type Controller struct {
serverSupportedSyncMessageVersion sharedgrpc.SyncMessageVersion
perAccountServerSupportedSyncMessageVersions map[string]sharedgrpc.SyncMessageVersion
nmdataStore networkmapdb.NetworkMapDBStore
}
type bufferUpdate struct {
@@ -83,7 +78,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, nmdataStore networkmapdb.NetworkMapDBStore) *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) *Controller {
nMetrics, err := newMetrics(metrics.UpdateChannelMetrics())
if err != nil {
log.Fatal(fmt.Errorf("error creating metrics: %w", err))
@@ -104,7 +99,6 @@ 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,
}
}
@@ -153,11 +147,6 @@ 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)
@@ -178,7 +167,7 @@ func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID strin
return nil
}
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
if err != nil {
return fmt.Errorf("failed to get validate peers: %v", err)
}
@@ -265,7 +254,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, types.TwinPeer(p), nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, types.TwinAccountSettings(account.Settings), extraSetting, maps.Keys(peerGroups), dnsFwdPort)
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)
c.metrics.CountToComponentSyncResponseDuration(time.Since(start))
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
@@ -286,7 +275,7 @@ func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID strin
}
start = time.Now()
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)
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)
c.metrics.CountToSyncResponseDuration(time.Since(start))
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
@@ -304,247 +293,6 @@ 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
@@ -577,10 +325,6 @@ 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)
@@ -596,7 +340,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, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
if err != nil {
return fmt.Errorf("failed to get validate peers: %v", err)
}
@@ -682,7 +426,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, types.TwinPeer(p), nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, types.TwinAccountSettings(account.Settings), extraSetting, maps.Keys(peerGroups), dnsFwdPort)
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)
c.metrics.CountToComponentSyncResponseDuration(time.Since(start))
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
@@ -703,7 +447,7 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s
}
start = time.Now()
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)
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)
c.metrics.CountToSyncResponseDuration(time.Since(start))
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
@@ -760,7 +504,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, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
if err != nil {
return fmt.Errorf("failed to get validated peers: %v", err)
}
@@ -820,7 +564,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, types.TwinPeer(peer), nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, types.TwinAccountSettings(account.Settings), extraSettings, maps.Keys(peerGroups), dnsFwdPort)
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)
c.peersUpdateManager.SendUpdate(ctx, peer.ID, &network_map.UpdateMessage{
Update: update,
@@ -837,7 +581,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, types.TwinPeer(peer), nil, nil, nmap, dnsDomain, postureChecks, dnsCache, types.TwinAccountSettings(account.Settings), extraSettings, maps.Keys(peerGroups), dnsFwdPort)
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)
c.peersUpdateManager.SendUpdate(ctx, peer.ID, &network_map.UpdateMessage{
Update: update,
@@ -897,7 +641,7 @@ func (c *Controller) GetValidatedPeerWithComponents(ctx context.Context, isRequi
if err != nil {
return nil, nil, nil, nil, 0, err
}
return peer, &types.NetworkMapComponents{Network: types.TwinNetwork(network)}, nil, nil, 0, nil
return peer, &types.NetworkMapComponents{Network: network.Copy()}, nil, nil, 0, nil
}
account, err := c.requestBuffer.GetAccountWithBackpressure(ctx, accountID)
@@ -907,7 +651,7 @@ func (c *Controller) GetValidatedPeerWithComponents(ctx context.Context, isRequi
c.injectAllProxyPolicies(ctx, account)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
if err != nil {
return nil, nil, nil, nil, 0, err
}
@@ -1050,7 +794,7 @@ func (c *Controller) GetValidatedPeerWithMap(ctx context.Context, isRequiresAppr
}
emptyMap := &types.NetworkMap{
Network: types.TwinNetwork(network),
Network: network.Copy(),
}
return emptyMap, nil, 0, nil
}
@@ -1062,7 +806,7 @@ func (c *Controller) GetValidatedPeerWithMap(ctx context.Context, isRequiresAppr
c.injectAllProxyPolicies(ctx, account)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
if err != nil {
return nil, nil, 0, err
}
@@ -1164,36 +908,20 @@ 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 {
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 {
if len(peers) == 0 {
return int64(network_map.OldForwarderPort)
}
reqVer := semver.Canonical(requiredVersion)
// Check if all peers have the required version or newer
for _, wtVersion := range wtVersions {
for _, peer := range peers {
// Development version is always supported
if version.IsDevelopmentVersion(wtVersion) {
if version.IsDevelopmentVersion(peer.Meta.WtVersion) {
continue
}
peerVersion := semver.Canonical("v" + wtVersion)
peerVersion := semver.Canonical("v" + peer.Meta.WtVersion)
if peerVersion == "" {
// If any peer doesn't have version info, return 0
return int64(network_map.OldForwarderPort)
@@ -1327,7 +1055,7 @@ func (c *Controller) GetNetworkMap(ctx context.Context, peerID string) (*types.N
groups[groupID] = group.Peers
}
validatedPeers, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
validatedPeers, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
if err != nil {
return nil, err
}

View File

@@ -420,47 +420,6 @@ var providers = []Provider{
{ID: "mistral-embed", Label: "Mistral Embed", InputPer1k: 0.0001, OutputPer1k: 0, ContextWindow: 8192},
},
},
{
ID: "kimi_api",
Kind: KindProvider,
Name: "Kimi (Moonshot AI) API",
Description: "Kimi K3 / K2 models via the Moonshot AI platform",
DefaultHost: "api.moonshot.ai",
AuthHeaderName: "Authorization",
AuthHeaderTemplate: "Bearer ${API_KEY}",
DefaultContentType: "application/json",
BrandColor: "#1A1A2E",
// ParserID empty on purpose: Moonshot serves two body shapes on
// the same host and key, and the proxy's URL sniffer dispatches
// both (same pattern as Bifrost). /v1/chat/completions matches
// OpenAIParser; the Anthropic-compatible endpoint the official
// Claude Code guide uses (/anthropic/v1/messages) contains
// "/v1/messages" and matches AnthropicParser. Pinning "openai"
// here would misparse the Claude Code path — the primary way
// teams consume Kimi for coding today. Both endpoints accept the
// same Moonshot key via Authorization: Bearer (Claude Code's
// ANTHROPIC_AUTH_TOKEN rides that header too).
//
// api.moonshot.ai is the international platform; mainland-China
// accounts live on api.moonshot.cn with separate billing —
// operators there override the host on the provider record. The
// kimi.com subscription coding endpoint (api.kimi.com/coding,
// model id "k3") is account-bound seat licensing rather than a
// meterable platform key, so it's deliberately not the default.
ParserID: "",
// Pricing per Moonshot's platform rates at K3 launch (July 2026):
// $3/$15 per MTok with $0.30 cached input, flat across the 1M-token
// window. kimi-k3 is the ONLY model the platform serves newer
// accounts — K2-era ids (kimi-k2-thinking) and even the kimi-latest
// alias return resource_not_found_error, verified live 2026-07-21 —
// so it's the only catalog entry. Grandfathered accounts with K2
// access can still type those ids on the provider's model rows.
// The consumer app's "K3 Swarm Max" mode is not an API SKU, so it
// doesn't appear here.
Models: []Model{
{ID: "kimi-k3", Label: "Kimi K3", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000},
},
},
{
ID: "litellm_proxy",
Kind: KindGateway,

View File

@@ -1,135 +0,0 @@
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])}
}

View File

@@ -1,84 +0,0 @@
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
}

View File

@@ -1,125 +0,0 @@
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,
}
}

View File

@@ -1,49 +0,0 @@
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
}

View File

@@ -1,38 +0,0 @@
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)
})
}
}

View File

@@ -1,72 +0,0 @@
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
}

View File

@@ -1,283 +0,0 @@
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)
}

View File

@@ -1,65 +0,0 @@
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
}

View File

@@ -1,57 +0,0 @@
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
}

View File

@@ -1,147 +0,0 @@
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
}

View File

@@ -1,65 +0,0 @@
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
}

View File

@@ -1,85 +0,0 @@
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
}

View File

@@ -1,52 +0,0 @@
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
}

View File

@@ -1,146 +0,0 @@
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"`
}

View File

@@ -1,56 +0,0 @@
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
}

View File

@@ -1,164 +0,0 @@
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"`
}

View File

@@ -1,56 +0,0 @@
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
}

View File

@@ -1,75 +0,0 @@
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
}

View File

@@ -1,219 +0,0 @@
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
}

View File

@@ -1,57 +0,0 @@
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
}

View File

@@ -7,7 +7,6 @@ import (
"crypto/tls"
"net/http"
"net/netip"
"os"
"slices"
"time"
@@ -21,16 +20,12 @@ import (
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/keepalive"
cachestore "github.com/eko/gocache/lib/v4/store"
"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"
@@ -73,8 +68,8 @@ func (s *BaseServer) Metrics() telemetry.AppMetrics {
// CacheStore returns a shared cache store backed by Redis or in-memory depending on the environment.
// All consumers should reuse this store to avoid creating multiple Redis connections.
func (s *BaseServer) CacheStore() cachestore.StoreInterface {
return Create(s, func() cachestore.StoreInterface {
func (s *BaseServer) CacheStore() nbcache.Store {
return Create(s, func() nbcache.Store {
cs, err := nbcache.NewStore(context.Background(), nbcache.DefaultStoreMaxTimeout, nbcache.DefaultStoreCleanupInterval, nbcache.DefaultStoreMaxConn)
if err != nil {
log.Fatalf("failed to create shared cache store: %v", err)
@@ -102,22 +97,6 @@ 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

View File

@@ -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, s.NetworkMapStore())
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)
})
}

View File

@@ -4,9 +4,13 @@ import (
"encoding/base64"
"strconv"
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/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"
)
@@ -81,7 +85,6 @@ 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,6 +106,7 @@ 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,
@@ -127,7 +131,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 *nmdata.Network) uint64 {
func networkSerial(n *types.Network) uint64 {
if n == nil {
return 0
}
@@ -140,15 +144,16 @@ type componentEncoder struct {
peerOrder map[string]uint32
peers []*proto.PeerCompact
networkIdToPublicId map[string]string
agentVersionOrder map[string]uint32
agentVersions []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)),
networkIdToPublicId: make(map[string]string),
components: c,
peerOrder: make(map[string]uint32, len(c.Peers)),
peers: make([]*proto.PeerCompact, 0, len(c.Peers)),
agentVersionOrder: make(map[string]uint32),
}
}
@@ -161,7 +166,7 @@ func (e *componentEncoder) indexAllPeers() {
}
}
func (e *componentEncoder) appendPeer(p *nmdata.Peer) uint32 {
func (e *componentEncoder) appendPeer(p *nbpeer.Peer) uint32 {
if idx, ok := e.peerOrder[p.ID]; ok {
return idx
}
@@ -175,7 +180,7 @@ func (e *componentEncoder) appendPeer(p *nmdata.Peer) 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]*nmdata.Peer) []uint32 {
func (e *componentEncoder) indexRouterPeers(routers map[string]*nbpeer.Peer) []uint32 {
if len(routers) == 0 {
return nil
}
@@ -189,15 +194,6 @@ func (e *componentEncoder) indexRouterPeers(routers map[string]*nmdata.Peer) []u
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
@@ -211,20 +207,10 @@ 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
@@ -234,7 +220,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 []*nmdata.Policy) []*proto.PolicyCompact {
func (e *componentEncoder) encodePolicies(policies []*types.Policy) []*proto.PolicyCompact {
if len(policies) == 0 {
return nil
}
@@ -256,7 +242,7 @@ func (e *componentEncoder) encodePolicies(policies []*nmdata.Policy) []*proto.Po
}
// encodePolicyRule maps a single PolicyRule under pol to a PolicyCompact entry.
func (e *componentEncoder) encodePolicyRule(pol *nmdata.Policy, r *nmdata.PolicyRule) *proto.PolicyCompact {
func (e *componentEncoder) encodePolicyRule(pol *types.Policy, r *types.PolicyRule) *proto.PolicyCompact {
return &proto.PolicyCompact{
Id: pol.PublicID,
Action: networkmap.GetProtoAction(string(r.Action)),
@@ -295,14 +281,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 []*nmdata.Policy, resourcePolicies map[string][]*nmdata.Policy) []*nmdata.Policy {
func unionPolicies(policies []*types.Policy, resourcePolicies map[string][]*types.Policy) []*types.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([]*nmdata.Policy, 0, len(policies))
out := make([]*types.Policy, 0, len(policies))
for _, p := range policies {
if p == nil {
continue
@@ -360,31 +346,18 @@ 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 nmdata.Resource) *proto.ResourceCompact {
t, ok := proto.ResourceCompactType_value[string(r.Type)]
if !ok || t == 0 || r.ID == "" {
func (e *componentEncoder) resourceToProto(r types.Resource) *proto.ResourceCompact {
if r.ID == "" && r.Type == "" {
return nil
}
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},
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
}
}
publicID, ok := e.networkIdToPublicId[r.ID]
if !ok {
return nil
}
return &proto.ResourceCompact{
Type: proto.ResourceCompactType(t),
ResourceId: &proto.ResourceCompact_Id{Id: publicID},
}
return out
}
// postureCheckSeqs translates a slice of posture-check xids to their
@@ -417,7 +390,7 @@ func (e *componentEncoder) networkPublicId(xid string) (string, bool) {
return id, true
}
func (e *componentEncoder) encodeDNSSettings(s *nmdata.DNSSettings) *proto.DNSSettingsCompact {
func (e *componentEncoder) encodeDNSSettings(s *types.DNSSettings) *proto.DNSSettingsCompact {
if s == nil || len(s.DisabledManagementGroups) == 0 {
return nil
}
@@ -432,7 +405,7 @@ func (e *componentEncoder) encodeDNSSettings(s *nmdata.DNSSettings) *proto.DNSSe
return out
}
func (e *componentEncoder) encodeRoutes(routes []*nmdata.Route) []*proto.RouteRaw {
func (e *componentEncoder) encodeRoutes(routes []*nbroute.Route) []*proto.RouteRaw {
if len(routes) == 0 {
return nil
}
@@ -470,7 +443,7 @@ func (e *componentEncoder) encodeRoutes(routes []*nmdata.Route) []*proto.RouteRa
return out
}
func (e *componentEncoder) encodeNameServerGroups(nsgs []*nmdata.NameServerGroup) []*proto.NameServerGroupRaw {
func (e *componentEncoder) encodeNameServerGroups(nsgs []*nbdns.NameServerGroup) []*proto.NameServerGroupRaw {
if len(nsgs) == 0 {
return nil
}
@@ -493,7 +466,7 @@ func (e *componentEncoder) encodeNameServerGroups(nsgs []*nmdata.NameServerGroup
return out
}
func encodeNameServers(servers []nmdata.NameServer) []*proto.NameServer {
func encodeNameServers(servers []nbdns.NameServer) []*proto.NameServer {
if len(servers) == 0 {
return nil
}
@@ -508,7 +481,7 @@ func encodeNameServers(servers []nmdata.NameServer) []*proto.NameServer {
return out
}
func encodeSimpleRecords(records []nmdata.SimpleRecord) []*proto.SimpleRecord {
func encodeSimpleRecords(records []nbdns.SimpleRecord) []*proto.SimpleRecord {
if len(records) == 0 {
return nil
}
@@ -525,7 +498,7 @@ func encodeSimpleRecords(records []nmdata.SimpleRecord) []*proto.SimpleRecord {
return out
}
func encodeCustomZones(zones []nmdata.CustomZone) []*proto.CustomZone {
func encodeCustomZones(zones []nbdns.CustomZone) []*proto.CustomZone {
if len(zones) == 0 {
return nil
}
@@ -541,7 +514,7 @@ func encodeCustomZones(zones []nmdata.CustomZone) []*proto.CustomZone {
return out
}
func (e *componentEncoder) encodeNetworkResources(resources []*nmdata.NetworkResource) []*proto.NetworkResourceRaw {
func (e *componentEncoder) encodeNetworkResources(resources []*resourceTypes.NetworkResource) []*proto.NetworkResourceRaw {
if len(resources) == 0 {
return nil
}
@@ -570,7 +543,7 @@ func (e *componentEncoder) encodeNetworkResources(resources []*nmdata.NetworkRes
return out
}
func (e *componentEncoder) encodeRoutersMap(routersMap map[string]map[string]*nmdata.NetworkRouter) map[string]*proto.NetworkRouterList {
func (e *componentEncoder) encodeRoutersMap(routersMap map[string]map[string]*routerTypes.NetworkRouter) map[string]*proto.NetworkRouterList {
if len(routersMap) == 0 {
return nil
}
@@ -606,7 +579,7 @@ func (e *componentEncoder) encodeRoutersMap(routersMap map[string]map[string]*nm
return out
}
func (e *componentEncoder) encodeResourcePoliciesMap(rpm map[string][]*nmdata.Policy) map[string]*proto.PolicyIds {
func (e *componentEncoder) encodeResourcePoliciesMap(rpm map[string][]*types.Policy) map[string]*proto.PolicyIds {
if len(rpm) == 0 {
return nil
}
@@ -693,7 +666,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 *nmdata.AccountSettingsInfo) *proto.AccountSettingsCompact {
func toAccountSettingsCompact(s *types.AccountSettingsInfo) *proto.AccountSettingsCompact {
if s == nil {
return &proto.AccountSettingsCompact{}
}
@@ -703,7 +676,7 @@ func toAccountSettingsCompact(s *nmdata.AccountSettingsInfo) *proto.AccountSetti
}
}
func toAccountNetwork(n *nmdata.Network) *proto.AccountNetwork {
func toAccountNetwork(n *types.Network) *proto.AccountNetwork {
if n == nil {
return nil
}
@@ -719,7 +692,7 @@ func toAccountNetwork(n *nmdata.Network) *proto.AccountNetwork {
return out
}
func toPeerCompact(p *nmdata.Peer) *proto.PeerCompact {
func toPeerCompact(p *nbpeer.Peer) *proto.PeerCompact {
pc := &proto.PeerCompact{
WgPubKey: decodeWgKey(p.Key),
SshPubKey: []byte(p.SSHKey),
@@ -781,7 +754,7 @@ func portsToUint32(ports []string) []uint32 {
return out
}
func portRangesToProto(ranges []nmdata.RulePortRange) []*proto.PortInfo_Range {
func portRangesToProto(ranges []types.RulePortRange) []*proto.PortInfo_Range {
if len(ranges) == 0 {
return nil
}

View File

@@ -15,8 +15,11 @@ import (
goproto "google.golang.org/protobuf/proto"
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/management/server/types"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
nbroute "github.com/netbirdio/netbird/route"
"github.com/netbirdio/netbird/shared/management/proto"
)
@@ -152,66 +155,67 @@ func envelopesEquivalent(a, b *proto.NetworkMapEnvelope) bool {
}
func newTestComponents() *types.NetworkMapComponents {
peerA := &nmdata.Peer{
peerA := &nbpeer.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"},
Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now()},
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
peerB := &nmdata.Peer{
peerB := &nbpeer.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"},
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.25.0"},
}
peerC := &nmdata.Peer{
peerC := &nbpeer.Peer{
ID: "peer-c",
Key: testWgKeyC,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 3}),
DNSLabel: "peerc",
Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
return &types.NetworkMapComponents{
PeerID: "peer-a",
Network: &nmdata.Network{
Network: &types.Network{
Identifier: "net-test",
Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)},
Serial: 7,
},
AccountSettings: &nmdata.AccountSettingsInfo{
AccountSettings: &types.AccountSettingsInfo{
PeerLoginExpirationEnabled: true,
PeerLoginExpiration: 2 * time.Hour,
},
Peers: map[string]*nmdata.Peer{
Peers: map[string]*nbpeer.Peer{
"peer-a": peerA,
"peer-b": peerB,
"peer-c": peerC,
},
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"}},
Groups: map[string]*types.Group{
"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"}},
},
Policies: []*nmdata.Policy{
Policies: []*types.Policy{
{
ID: "pol-1",
PublicID: "10",
Enabled: true,
Rules: []*nmdata.PolicyRule{{
ID: "rule-1", Enabled: true, Action: string(types.PolicyTrafficActionAccept),
Protocol: string(types.PolicyRuleProtocolTCP), Bidirectional: true,
Rules: []*types.PolicyRule{{
ID: "rule-1", Enabled: true, Action: types.PolicyTrafficActionAccept,
Protocol: types.PolicyRuleProtocolTCP, Bidirectional: true,
Ports: []string{"22", "80"},
PortRanges: []nmdata.RulePortRange{{Start: 8000, End: 8100}},
PortRanges: []types.RulePortRange{{Start: 8000, End: 8100}},
Sources: []string{"group-src"},
Destinations: []string{"group-dst"},
}},
},
},
RouterPeers: map[string]*nmdata.Peer{"peer-c": peerC},
RouterPeers: map[string]*nbpeer.Peer{"peer-c": peerC},
}
}
@@ -304,31 +308,6 @@ 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()
@@ -402,12 +381,12 @@ func TestEncodeNetworkMapEnvelope_MalformedWgKey(t *testing.T) {
func TestEncodeNetworkMapEnvelope_IPv6OnlyPeer(t *testing.T) {
c := newTestComponents()
v6Only := &nmdata.Peer{
v6Only := &nbpeer.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"},
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
c.Peers["peer-v6"] = v6Only
@@ -426,11 +405,11 @@ func TestEncodeNetworkMapEnvelope_IPv6OnlyPeer(t *testing.T) {
func TestEncodeNetworkMapEnvelope_PeerWithoutIP(t *testing.T) {
c := newTestComponents()
c.Peers["peer-noip"] = &nmdata.Peer{
c.Peers["peer-noip"] = &nbpeer.Peer{
ID: "peer-noip",
Key: testWgKeyA,
DNSLabel: "peernoip",
Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
full := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c}).GetFull()
@@ -448,7 +427,7 @@ func TestEncodeNetworkMapEnvelope_PeerWithoutIP(t *testing.T) {
func TestEncodeNetworkMapEnvelope_EmptyInput(t *testing.T) {
c := &types.NetworkMapComponents{
Network: &nmdata.Network{Identifier: "x", Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}},
Network: &types.Network{Identifier: "x", Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}},
}
env := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c})
@@ -497,7 +476,7 @@ func TestEncodeNetworkMapEnvelope_PeerLoginExpirationFields(t *testing.T) {
func TestEncodeNetworkMapEnvelope_RoutesRoundTrip(t *testing.T) {
c := newTestComponents()
c.Routes = []*nmdata.Route{
c.Routes = []*nbroute.Route{
{
ID: "route-peer",
PublicID: "100",
@@ -544,7 +523,7 @@ func TestEncodeNetworkMapEnvelope_RoutesRoundTrip(t *testing.T) {
func TestEncodeNetworkMapEnvelope_RouteWithMissingPeerLeavesIndexUnset(t *testing.T) {
c := newTestComponents()
c.Routes = []*nmdata.Route{{
c.Routes = []*nbroute.Route{{
ID: "route-x",
PublicID: "100",
Peer: "peer-not-in-components",
@@ -564,21 +543,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 := &nmdata.Policy{
resourceOnlyPolicy := &types.Policy{
ID: "pol-resource", PublicID: "99", Enabled: true,
Rules: []*nmdata.PolicyRule{{
ID: "rule-r", Enabled: true, Action: string(types.PolicyTrafficActionAccept),
Protocol: string(types.PolicyRuleProtocolTCP),
Rules: []*types.PolicyRule{{
ID: "rule-r", Enabled: true, Action: types.PolicyTrafficActionAccept,
Protocol: types.PolicyRuleProtocolTCP,
Sources: []string{"group-src"},
Destinations: []string{"group-dst"},
}},
}
c.ResourcePoliciesMap = map[string][]*nmdata.Policy{
c.ResourcePoliciesMap = map[string][]*types.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 = []*nmdata.NetworkResource{
c.NetworkResources = []*resourceTypes.NetworkResource{
{ID: "resource-x", PublicID: "77", Name: "res-x", Enabled: true},
}
@@ -587,16 +566,27 @@ 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 = []*nmdata.NameServerGroup{{
c.NameServerGroups = []*nbdns.NameServerGroup{{
ID: "nsg-1", PublicID: "50", Name: "Main", Description: "primary",
NameServers: []nmdata.NameServer{{
IP: netip.MustParseAddr("8.8.8.8"), NSType: int(nbdns.UDPNameServerType), Port: 53,
NameServers: []nbdns.NameServer{{
IP: netip.MustParseAddr("8.8.8.8"), NSType: nbdns.UDPNameServerType, Port: 53,
}},
Groups: []string{"group-src", "group-not-persisted"},
Primary: true, Enabled: true,
@@ -635,11 +625,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]*nmdata.NetworkRouter{
c.RoutersMap = map[string]map[string]*routerTypes.NetworkRouter{
"net-1": {
"peer-c": {
PublicID: "200",
Masquerade: true, Metric: 10, Enabled: true,
ID: "router-1", PublicID: "200",
Peer: "peer-c", Masquerade: true, Metric: 10, Enabled: true,
},
},
}
@@ -665,14 +655,14 @@ func TestEncodeNetworkMapEnvelope_RouterPeerNotInComponentsPeers(t *testing.T) {
// peer_index reference must still resolve.
c := newTestComponents()
delete(c.Peers, "peer-c")
routerPeer := &nmdata.Peer{
routerPeer := &nbpeer.Peer{
ID: "peer-c", Key: testWgKeyC, IP: netip.AddrFrom4([4]byte{100, 64, 0, 3}),
DNSLabel: "peerc", Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
DNSLabel: "peerc", Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
c.RouterPeers = map[string]*nmdata.Peer{"peer-c": routerPeer}
c.RouterPeers = map[string]*nbpeer.Peer{"peer-c": routerPeer}
c.NetworkXIDToPublicID = map[string]string{"net-1": "5"}
c.RoutersMap = map[string]map[string]*nmdata.NetworkRouter{
"net-1": {"peer-c": {PublicID: "1", Enabled: true}},
c.RoutersMap = map[string]map[string]*routerTypes.NetworkRouter{
"net-1": {"peer-c": {ID: "r-1", PublicID: "1", Peer: "peer-c", Enabled: true}},
}
full := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c}).GetFull()
@@ -705,9 +695,9 @@ func TestToProxyPatch_EmptyInputReturnsNil(t *testing.T) {
func TestToProxyPatch_PopulatesAllFields(t *testing.T) {
nm := &types.NetworkMap{
Peers: []*nmdata.Peer{{
Peers: []*nbpeer.Peer{{
ID: "ext-peer", Key: testWgKeyA, IP: netip.AddrFrom4([4]byte{100, 64, 0, 9}),
DNSLabel: "extpeer", Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
DNSLabel: "extpeer", Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}},
FirewallRules: []*types.FirewallRule{{
PeerIP: "100.64.0.9", Action: "accept", Direction: 0, Protocol: "tcp",
@@ -776,7 +766,7 @@ func TestEncodeNetworkMapEnvelope_NilComponentsGracefulDegrade(t *testing.T) {
func TestEncodeNetworkMapEnvelope_AccountSettingsAlwaysEmitted(t *testing.T) {
c := &types.NetworkMapComponents{
Network: &nmdata.Network{Identifier: "x", Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}},
Network: &types.Network{Identifier: "x", Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}},
// AccountSettings deliberately nil
}
@@ -790,6 +780,6 @@ func TestEncodeNetworkMapEnvelope_AccountSettingsAlwaysEmitted(t *testing.T) {
func emptyNetworkMapComponents() *types.NetworkMapComponents {
return types.EmptyNetworkMapComponents(
&types.NetworkMapComponents{
PeerID: "peer-id", Peers: map[string]*nmdata.Peer{"peer-id": {}}},
PeerID: "peer-id", Peers: map[string]*nbpeer.Peer{"peer-id": {}}},
)
}

View File

@@ -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 *nmdata.Peer,
peer *nbpeer.Peer,
turnCredentials *Token,
relayCredentials *Token,
components *types.NetworkMapComponents,
proxyPatch *types.NetworkMap,
dnsName string,
checks []*posture.Checks,
settings *nmdata.AccountSettingsInfo,
settings *types.Settings,
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, components.ForceRoutingPeerDNSResolution)
peerConfig := toPeerConfig(peer, components.Network, dnsName, settings, httpConfig, deviceFlowConfig, enableSSH)
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 *nmdata.Peer) bool {
func computeSSHEnabledForPeer(c *types.NetworkMapComponents, peer *nbpeer.Peer) bool {
if c == nil || peer == nil {
return false
}
@@ -170,25 +170,25 @@ func computeSSHEnabledForPeer(c *types.NetworkMapComponents, peer *nmdata.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 *nmdata.PolicyRule, peer *nmdata.Peer) bool {
func ruleEnablesSSHForPeer(c *types.NetworkMapComponents, rule *types.PolicyRule, peer *nbpeer.Peer) bool {
if rule == nil || !rule.Enabled {
return false
}
if !peerInDestinations(c, rule, peer.ID) {
return false
}
if rule.Protocol == string(types.PolicyRuleProtocolNetbirdSSH) {
if rule.Protocol == types.PolicyRuleProtocolNetbirdSSH {
return true
}
return peer.SSHEnabled && nmdata.PolicyRuleImpliesLegacySSH(rule)
return peer.SSHEnabled && types.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 *nmdata.PolicyRule, peerID string) bool {
if rule.DestinationResource.Type == string(types.ResourceTypePeer) && rule.DestinationResource.ID != "" {
func peerInDestinations(c *types.NetworkMapComponents, rule *types.PolicyRule, peerID string) bool {
if rule.DestinationResource.Type == types.ResourceTypePeer && rule.DestinationResource.ID != "" {
return rule.DestinationResource.ID == peerID
}
for _, groupID := range rule.Destinations {

View File

@@ -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,15 +17,16 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
const targetPeerID = "target"
const targetGroupID = "g_dst"
mkComponents := func(rule *nmdata.PolicyRule, sshEnabled bool) (*types.NetworkMapComponents, *nmdata.Peer) {
peer := &nmdata.Peer{ID: targetPeerID, SSHEnabled: sshEnabled}
mkComponents := func(rule *types.PolicyRule, sshEnabled bool) (*types.NetworkMapComponents, *nbpeer.Peer) {
peer := &nbpeer.Peer{ID: targetPeerID, SSHEnabled: sshEnabled}
group := &types.Group{ID: targetGroupID, Name: "dst", Peers: []string{targetPeerID}}
return &types.NetworkMapComponents{
Peers: map[string]*nmdata.Peer{targetPeerID: peer},
Groups: map[string]*nmdata.Group{targetGroupID: {Name: "dst", Peers: []string{targetPeerID}}},
Policies: []*nmdata.Policy{{
Peers: map[string]*nbpeer.Peer{targetPeerID: peer},
Groups: map[string]*types.Group{targetGroupID: group},
Policies: []*types.Policy{{
ID: "p",
Enabled: true,
Rules: []*nmdata.PolicyRule{rule},
Rules: []*types.PolicyRule{rule},
}},
}, peer
}
@@ -33,14 +34,14 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
cases := []struct {
name string
peerSSH bool
rule nmdata.PolicyRule
rule types.PolicyRule
wantEnabled bool
}{
{
name: "explicit-netbird-ssh-activates-regardless-of-peer-ssh",
peerSSH: false,
rule: nmdata.PolicyRule{
Enabled: true, Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
rule: types.PolicyRule{
Enabled: true, Protocol: types.PolicyRuleProtocolNetbirdSSH,
Destinations: []string{targetGroupID},
},
wantEnabled: true,
@@ -48,8 +49,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "implicit-tcp-22-with-peer-ssh",
peerSSH: true,
rule: nmdata.PolicyRule{
Enabled: true, Protocol: string(types.PolicyRuleProtocolTCP), Ports: []string{"22"},
rule: types.PolicyRule{
Enabled: true, Protocol: types.PolicyRuleProtocolTCP, Ports: []string{"22"},
Destinations: []string{targetGroupID},
},
wantEnabled: true,
@@ -57,8 +58,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "implicit-tcp-22-without-peer-ssh-disabled",
peerSSH: false,
rule: nmdata.PolicyRule{
Enabled: true, Protocol: string(types.PolicyRuleProtocolTCP), Ports: []string{"22"},
rule: types.PolicyRule{
Enabled: true, Protocol: types.PolicyRuleProtocolTCP, Ports: []string{"22"},
Destinations: []string{targetGroupID},
},
wantEnabled: false,
@@ -66,8 +67,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "implicit-tcp-22022-with-peer-ssh",
peerSSH: true,
rule: nmdata.PolicyRule{
Enabled: true, Protocol: string(types.PolicyRuleProtocolTCP), Ports: []string{"22022"},
rule: types.PolicyRule{
Enabled: true, Protocol: types.PolicyRuleProtocolTCP, Ports: []string{"22022"},
Destinations: []string{targetGroupID},
},
wantEnabled: true,
@@ -75,8 +76,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "implicit-all-protocol-with-peer-ssh",
peerSSH: true,
rule: nmdata.PolicyRule{
Enabled: true, Protocol: string(types.PolicyRuleProtocolALL),
rule: types.PolicyRule{
Enabled: true, Protocol: types.PolicyRuleProtocolALL,
Destinations: []string{targetGroupID},
},
wantEnabled: true,
@@ -84,10 +85,10 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "implicit-port-range-covers-22",
peerSSH: true,
rule: nmdata.PolicyRule{
rule: types.PolicyRule{
Enabled: true,
Protocol: string(types.PolicyRuleProtocolTCP),
PortRanges: []nmdata.RulePortRange{{Start: 20, End: 30}},
Protocol: types.PolicyRuleProtocolTCP,
PortRanges: []types.RulePortRange{{Start: 20, End: 30}},
Destinations: []string{targetGroupID},
},
wantEnabled: true,
@@ -95,8 +96,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "tcp-80-no-ssh",
peerSSH: true,
rule: nmdata.PolicyRule{
Enabled: true, Protocol: string(types.PolicyRuleProtocolTCP), Ports: []string{"80"},
rule: types.PolicyRule{
Enabled: true, Protocol: types.PolicyRuleProtocolTCP, Ports: []string{"80"},
Destinations: []string{targetGroupID},
},
wantEnabled: false,
@@ -104,8 +105,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "disabled-rule-skipped",
peerSSH: true,
rule: nmdata.PolicyRule{
Enabled: false, Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
rule: types.PolicyRule{
Enabled: false, Protocol: types.PolicyRuleProtocolNetbirdSSH,
Destinations: []string{targetGroupID},
},
wantEnabled: false,
@@ -113,8 +114,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "peer-not-in-destinations",
peerSSH: true,
rule: nmdata.PolicyRule{
Enabled: true, Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
rule: types.PolicyRule{
Enabled: true, Protocol: types.PolicyRuleProtocolNetbirdSSH,
Destinations: []string{"g_other"}, // target not in this group
},
wantEnabled: false,
@@ -122,21 +123,21 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "peer-typed-destination-resource-matches",
peerSSH: false,
rule: nmdata.PolicyRule{
rule: types.PolicyRule{
Enabled: true,
Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
DestinationResource: nmdata.Resource{ID: targetPeerID, Type: string(types.ResourceTypePeer)},
Protocol: types.PolicyRuleProtocolNetbirdSSH,
DestinationResource: types.Resource{ID: targetPeerID, Type: types.ResourceTypePeer},
},
wantEnabled: true,
},
{
name: "non-peer-destination-resource-falls-through-to-groups",
peerSSH: false,
rule: nmdata.PolicyRule{
rule: types.PolicyRule{
Enabled: true,
Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
DestinationResource: nmdata.Resource{ID: targetPeerID, Type: "host"}, // wrong type
Destinations: []string{targetGroupID}, // saved by group fallback
Protocol: types.PolicyRuleProtocolNetbirdSSH,
DestinationResource: types.Resource{ID: targetPeerID, Type: "host"}, // wrong type
Destinations: []string{targetGroupID}, // saved by group fallback
},
wantEnabled: true,
},
@@ -155,16 +156,16 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
// belt-and-suspenders presence guard mirroring Calculate's
// getAllPeersFromGroups invariant.
func TestComputeSSHEnabledForPeer_TargetMissingFromComponents(t *testing.T) {
peer := &nmdata.Peer{ID: "missing", SSHEnabled: true}
peer := &nbpeer.Peer{ID: "missing", SSHEnabled: true}
c := &types.NetworkMapComponents{
Peers: map[string]*nmdata.Peer{}, // target peer NOT present
Groups: map[string]*nmdata.Group{
"g": {Peers: []string{"missing"}},
Peers: map[string]*nbpeer.Peer{}, // target peer NOT present
Groups: map[string]*types.Group{
"g": {ID: "g", Peers: []string{"missing"}},
},
Policies: []*nmdata.Policy{{
Policies: []*types.Policy{{
ID: "p", Enabled: true,
Rules: []*nmdata.PolicyRule{{
Enabled: true, Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
Rules: []*types.PolicyRule{{
Enabled: true, Protocol: types.PolicyRuleProtocolNetbirdSSH,
Destinations: []string{"g"},
}},
}},
@@ -178,6 +179,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, &nmdata.Peer{ID: "x"}))
assert.False(t, computeSSHEnabledForPeer(nil, &nbpeer.Peer{ID: "x"}))
assert.False(t, computeSSHEnabledForPeer(&types.NetworkMapComponents{}, nil))
}

View File

@@ -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 *nmdata.AccountSettingsInfo) *proto.NetbirdConfig {
func toNetbirdConfig(config *nbconfig.Config, turnCredentials *Token, relayToken *Token, extraSettings *types.ExtraSettings, settings *types.Settings) *proto.NetbirdConfig {
if config == nil {
return nil
}
@@ -119,7 +119,7 @@ func toNetbirdConfig(config *nbconfig.Config, turnCredentials *Token, relayToken
return nbConfig
}
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 {
func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, settings *types.Settings, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, enableSSH bool) *proto.PeerConfig {
netmask, _ := network.Net.Mask.Size()
fqdn := peer.FQDN(dnsName)
@@ -135,7 +135,7 @@ func toPeerConfig(peer *nmdata.Peer, network *nmdata.Network, dnsName string, se
Address: fmt.Sprintf("%s/%d", peer.IP.String(), netmask),
SshConfig: sshConfig,
Fqdn: fqdn,
RoutingPeerDnsResolutionEnabled: settings.RoutingPeerDNSResolutionEnabled || peer.ProxyMeta.Embedded || forceRoutingPeerDNS,
RoutingPeerDnsResolutionEnabled: settings.RoutingPeerDNSResolutionEnabled,
LazyConnectionEnabled: settings.LazyConnectionEnabled,
AutoUpdate: &proto.AutoUpdateSettings{
Version: settings.AutoUpdateVersion,
@@ -154,7 +154,7 @@ func toPeerConfig(peer *nmdata.Peer, network *nmdata.Network, dnsName string, se
return peerConfig
}
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 {
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 {
// 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, networkMap.ForceRoutingPeerDNSResolution),
PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH),
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, networkMap.ForceRoutingPeerDNSResolution),
PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH),
},
Checks: toProtocolChecks(ctx, checks),
}

View File

@@ -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, types.TwinAccountSettings(settings))
nbCfg := toNetbirdConfig(nil, nil, nil, nil, 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, types.TwinAccountSettings(settings))
nbCfg := toNetbirdConfig(cfg, nil, relayToken, nil, 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")

View File

@@ -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, 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),
NetbirdConfig: toNetbirdConfig(s.config, nil, relayToken, nil, settings),
PeerConfig: toPeerConfig(peer, network, s.networkMapController.GetDNSDomain(settings), settings, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, enableSSH),
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, types.TwinPeer(freshPeer), turnToken, relayToken, components, proxyPatch, dnsName, freshPostureChecks, types.TwinAccountSettings(settings), settings.Extra, peerGroups, freshDnsFwdPort)
plainResp = ToComponentSyncResponse(ctx, s.config, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, freshPeer, turnToken, relayToken, components, proxyPatch, dnsName, freshPostureChecks, settings, settings.Extra, peerGroups, freshDnsFwdPort)
} else {
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)
plainResp = ToSyncResponse(ctx, s.config, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, peer, turnToken, relayToken, networkMap, dnsName, postureChecks, nil, settings, settings.Extra, peerGroups, dnsFwdPort)
}
key, err := s.secretsManager.GetWGKey()

View File

@@ -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{}, nil)
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{})
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

View File

@@ -7,9 +7,6 @@ import (
"errors"
"fmt"
"time"
"github.com/eko/gocache/lib/v4/cache"
"github.com/eko/gocache/lib/v4/store"
)
const (
@@ -22,12 +19,17 @@ var (
ErrTokenExpired = errors.New("JWT expired")
)
type SessionStore struct {
cache *cache.Cache[string]
// TokenCache atomically records used JWTs until their expiration.
type TokenCache interface {
SetNX(ctx context.Context, key, value string, ttl time.Duration) (bool, error)
}
func NewSessionStore(cacheStore store.StoreInterface) *SessionStore {
return &SessionStore{cache: cache.New[string](cacheStore)}
type SessionStore struct {
cache TokenCache
}
func NewSessionStore(cacheStore TokenCache) *SessionStore {
return &SessionStore{cache: cacheStore}
}
// RegisterToken records a JWT until its exp time and rejects reuse.
@@ -38,19 +40,13 @@ func (s *SessionStore) RegisterToken(ctx context.Context, token string, expiresA
}
key := usedTokenKeyPrefix + hashToken(token)
_, err := s.cache.Get(ctx, key)
if err == nil {
return ErrTokenAlreadyUsed
}
var notFound *store.NotFound
if !errors.As(err, &notFound) {
return fmt.Errorf("failed to lookup used token entry: %w", err)
}
if err := s.cache.Set(ctx, key, usedTokenMarker, store.WithExpiration(ttl)); err != nil {
created, err := s.cache.SetNX(ctx, key, usedTokenMarker, ttl)
if err != nil {
return fmt.Errorf("failed to store used token entry: %w", err)
}
if !created {
return ErrTokenAlreadyUsed
}
return nil
}

View File

@@ -2,6 +2,7 @@ package auth
import (
"context"
"errors"
"testing"
"time"
@@ -38,6 +39,39 @@ func TestSessionStore_RegisterSameTokenTwiceIsRejected(t *testing.T) {
assert.ErrorIs(t, err, ErrTokenAlreadyUsed)
}
func TestSessionStore_ConcurrentRegistrationAllowsOneCaller(t *testing.T) {
s := newTestSessionStore(t)
ctx := context.Background()
const attempts = 100
start := make(chan struct{})
results := make(chan error, attempts)
for range attempts {
go func() {
<-start
results <- s.RegisterToken(ctx, "token", time.Now().Add(time.Hour))
}()
}
close(start)
succeeded := 0
alreadyUsed := 0
for range attempts {
err := <-results
switch {
case err == nil:
succeeded++
case errors.Is(err, ErrTokenAlreadyUsed):
alreadyUsed++
default:
require.NoError(t, err)
}
}
assert.Equal(t, 1, succeeded)
assert.Equal(t, attempts-1, alreadyUsed)
}
func TestSessionStore_RegisterDifferentTokensAreIndependent(t *testing.T) {
s := newTestSessionStore(t)
ctx := context.Background()
@@ -72,6 +106,23 @@ func TestSessionStore_EntryEvictsAtTTLAndAllowsReRegistration(t *testing.T) {
require.NoError(t, s.RegisterToken(ctx, token, time.Now().Add(time.Hour)))
}
type failingTokenCache struct {
err error
}
func (f failingTokenCache) SetNX(context.Context, string, string, time.Duration) (bool, error) {
return false, f.err
}
func TestSessionStore_CacheErrorIsReturned(t *testing.T) {
cacheErr := errors.New("cache unavailable")
s := NewSessionStore(failingTokenCache{err: cacheErr})
err := s.RegisterToken(context.Background(), "token", time.Now().Add(time.Hour))
require.Error(t, err)
assert.ErrorIs(t, err, cacheErr)
}
func TestHashToken_StableAndDoesNotLeak(t *testing.T) {
a := hashToken("tokenA")
b := hashToken("tokenB")

31
management/server/cache/memory.go vendored Normal file
View File

@@ -0,0 +1,31 @@
package cache
import (
"context"
"time"
"github.com/eko/gocache/lib/v4/store"
gocachestore "github.com/eko/gocache/store/go_cache/v4"
gocache "github.com/patrickmn/go-cache"
)
type goCacheStore struct {
store.StoreInterface
client *gocache.Cache
}
func newMemoryStore(maxTimeout, cleanupInterval time.Duration) Store {
client := gocache.New(maxTimeout, cleanupInterval)
return &goCacheStore{
StoreInterface: gocachestore.NewGoCache(client),
client: client,
}
}
func (s *goCacheStore) SetNX(_ context.Context, key, value string, ttl time.Duration) (bool, error) {
// Add only returns an error when a non-expired entry already exists.
if err := s.client.Add(key, value, ttl); err != nil {
return false, nil
}
return true, nil
}

49
management/server/cache/memory_test.go vendored Normal file
View File

@@ -0,0 +1,49 @@
package cache_test
import (
"context"
"testing"
"time"
"github.com/netbirdio/netbird/management/server/cache"
)
func TestMemoryStore(t *testing.T) {
memStore, err := cache.NewStore(context.Background(), 100*time.Millisecond, 300*time.Millisecond, 100)
if err != nil {
t.Fatalf("couldn't create memory store: %s", err)
}
ctx := context.Background()
key, value := "testing", "tested"
err = memStore.Set(ctx, key, value)
if err != nil {
t.Errorf("couldn't set testing data: %s", err)
}
result, err := memStore.Get(ctx, key)
if err != nil {
t.Errorf("couldn't get testing data: %s", err)
}
if value != result.(string) {
t.Errorf("value returned doesn't match testing data, got %s, expected %s", result, value)
}
created, err := memStore.SetNX(ctx, "atomic", value, 100*time.Millisecond)
if err != nil {
t.Fatalf("couldn't atomically set testing data: %s", err)
}
if !created {
t.Fatal("first atomic set should create the entry")
}
created, err = memStore.SetNX(ctx, "atomic", value, 100*time.Millisecond)
if err != nil {
t.Fatalf("couldn't atomically check testing data: %s", err)
}
if created {
t.Fatal("second atomic set should not replace the entry")
}
// test expiration
time.Sleep(300 * time.Millisecond)
_, err = memStore.Get(ctx, key)
if err == nil {
t.Error("value should not be found")
}
}

51
management/server/cache/redis.go vendored Normal file
View File

@@ -0,0 +1,51 @@
package cache
import (
"context"
"fmt"
"math"
"time"
"github.com/eko/gocache/lib/v4/store"
redisstore "github.com/eko/gocache/store/redis/v4"
"github.com/redis/go-redis/v9"
log "github.com/sirupsen/logrus"
)
type redisStore struct {
store.StoreInterface
client *redis.Client
}
func getRedisStore(ctx context.Context, redisEnvAddr string, maxConn int) (Store, error) {
options, err := redis.ParseURL(redisEnvAddr)
if err != nil {
return nil, fmt.Errorf("parsing redis cache url: %s", err)
}
options.MaxIdleConns = int(math.Ceil(float64(maxConn) * 0.5)) // 50% of max conns
options.MinIdleConns = int(math.Ceil(float64(maxConn) * 0.1)) // 10% of max conns
options.MaxActiveConns = maxConn
options.ConnMaxIdleTime = 30 * time.Minute
options.ConnMaxLifetime = 0
options.PoolTimeout = 10 * time.Second
redisClient := redis.NewClient(options)
subCtx, cancel := context.WithTimeout(ctx, 2*time.Second)
defer cancel()
_, err = redisClient.Ping(subCtx).Result()
if err != nil {
return nil, err
}
log.WithContext(subCtx).Infof("using redis cache at %s", redisEnvAddr)
return &redisStore{
StoreInterface: redisstore.NewRedis(redisClient),
client: redisClient,
}, nil
}
func (s *redisStore) SetNX(ctx context.Context, key, value string, ttl time.Duration) (bool, error) {
return s.client.SetNX(ctx, key, value, ttl).Result()
}

View File

@@ -12,32 +12,6 @@ import (
"github.com/netbirdio/netbird/management/server/cache"
)
func TestMemoryStore(t *testing.T) {
memStore, err := cache.NewStore(context.Background(), 100*time.Millisecond, 300*time.Millisecond, 100)
if err != nil {
t.Fatalf("couldn't create memory store: %s", err)
}
ctx := context.Background()
key, value := "testing", "tested"
err = memStore.Set(ctx, key, value)
if err != nil {
t.Errorf("couldn't set testing data: %s", err)
}
result, err := memStore.Get(ctx, key)
if err != nil {
t.Errorf("couldn't get testing data: %s", err)
}
if value != result.(string) {
t.Errorf("value returned doesn't match testing data, got %s, expected %s", result, value)
}
// test expiration
time.Sleep(300 * time.Millisecond)
_, err = memStore.Get(ctx, key)
if err == nil {
t.Error("value should not be found")
}
}
func TestRedisStoreConnectionFailure(t *testing.T) {
t.Setenv(cache.RedisStoreEnvVar, "redis://127.0.0.1:6379")
_, err := cache.NewStore(context.Background(), 10*time.Millisecond, 30*time.Millisecond, 100)
@@ -94,6 +68,47 @@ func TestRedisStoreConnectionSuccess(t *testing.T) {
if value != r {
t.Errorf("value returned from redis doesn't match testing data, got %s, expected %s", r, value)
}
secondRedisStore, err := cache.NewStore(context.Background(), 100*time.Millisecond, 300*time.Millisecond, 100)
if err != nil {
t.Fatalf("couldn't create second redis store: %s", err)
}
start := make(chan struct{})
type setResult struct {
created bool
err error
}
results := make(chan setResult, 2)
for _, cacheStore := range []cache.Store{redisStore, secondRedisStore} {
go func() {
<-start
created, err := cacheStore.SetNX(ctx, "atomic", value, time.Second)
results <- setResult{created: created, err: err}
}()
}
close(start)
created := 0
for range 2 {
result := <-results
if result.err != nil {
t.Fatalf("atomic redis set failed: %s", result.err)
}
if result.created {
created++
}
}
if created != 1 {
t.Fatalf("expected exactly one redis client to create the entry, got %d", created)
}
ttl, err := redisClient.PTTL(ctx, "atomic").Result()
if err != nil {
t.Fatalf("couldn't read atomic entry TTL: %s", err)
}
if ttl <= 0 {
t.Fatalf("atomic entry should have a positive TTL, got %s", ttl)
}
// test expiration
time.Sleep(300 * time.Millisecond)
_, err = redisStore.Get(ctx, key)

View File

@@ -2,17 +2,10 @@ package cache
import (
"context"
"fmt"
"math"
"os"
"time"
"github.com/eko/gocache/lib/v4/store"
gocache_store "github.com/eko/gocache/store/go_cache/v4"
redis_store "github.com/eko/gocache/store/redis/v4"
gocache "github.com/patrickmn/go-cache"
"github.com/redis/go-redis/v9"
log "github.com/sirupsen/logrus"
)
// RedisStoreEnvVar is the environment variable that determines if a redis store should be used.
@@ -31,15 +24,21 @@ const (
DefaultStoreMaxConn = 1000
)
// Store extends the shared cache interface with atomic insertion support.
type Store interface {
store.StoreInterface
// SetNX atomically stores a value with a TTL only when the key does not exist.
SetNX(ctx context.Context, key, value string, ttl time.Duration) (bool, error)
}
// NewStore creates a new cache store with the given max timeout and cleanup interval. It checks for the environment Variable RedisStoreEnvVar
// to determine if a redis store should be used. If the environment variable is set, it will attempt to connect to the redis store.
func NewStore(ctx context.Context, maxTimeout, cleanupInterval time.Duration, maxConn int) (store.StoreInterface, error) {
func NewStore(ctx context.Context, maxTimeout, cleanupInterval time.Duration, maxConn int) (Store, error) {
redisAddr := GetAddrFromEnv()
if redisAddr != "" {
return getRedisStore(ctx, redisAddr, maxConn)
}
goc := gocache.New(maxTimeout, cleanupInterval)
return gocache_store.NewGoCache(goc), nil
return newMemoryStore(maxTimeout, cleanupInterval), nil
}
// GetAddrFromEnv returns the redis address from the environment variable RedisStoreEnvVar or its legacy counterpart.
@@ -50,29 +49,3 @@ func GetAddrFromEnv() string {
}
return addr
}
func getRedisStore(ctx context.Context, redisEnvAddr string, maxConn int) (store.StoreInterface, error) {
options, err := redis.ParseURL(redisEnvAddr)
if err != nil {
return nil, fmt.Errorf("parsing redis cache url: %s", err)
}
options.MaxIdleConns = int(math.Ceil(float64(maxConn) * 0.5)) // 50% of max conns
options.MinIdleConns = int(math.Ceil(float64(maxConn) * 0.1)) // 10% of max conns
options.MaxActiveConns = maxConn
options.ConnMaxIdleTime = 30 * time.Minute
options.ConnMaxLifetime = 0
options.PoolTimeout = 10 * time.Second
redisClient := redis.NewClient(options)
subCtx, cancel := context.WithTimeout(ctx, 2*time.Second)
defer cancel()
_, err = redisClient.Ping(subCtx).Result()
if err != nil {
return nil, err
}
log.WithContext(subCtx).Infof("using redis cache at %s", redisEnvAddr)
return redis_store.NewRedis(redisClient), nil
}

View File

@@ -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{}, nil)
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{})
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)
}

View File

@@ -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(account.Peers, netMap, dnsDomain))
util.WriteJSONObject(ctx, w, toAccessiblePeers(netMap, dnsDomain))
}
func (h *Handler) CreateTemporaryAccess(w http.ResponseWriter, r *http.Request) {
@@ -534,21 +534,14 @@ func (h *Handler) CreateTemporaryAccess(w http.ResponseWriter, r *http.Request)
util.WriteJSONObject(r.Context(), w, resp)
}
// 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 {
func toAccessiblePeers(netMap *types.NetworkMap, dnsDomain string) []api.AccessiblePeer {
accessiblePeers := make([]api.AccessiblePeer, 0, len(netMap.Peers)+len(netMap.OfflinePeers))
appendByID := func(id string) {
if p, ok := accountPeers[id]; ok && p != nil {
accessiblePeers = append(accessiblePeers, peerToAccessiblePeer(p, dnsDomain))
}
}
for _, p := range netMap.Peers {
appendByID(p.ID)
accessiblePeers = append(accessiblePeers, peerToAccessiblePeer(p, dnsDomain))
}
for _, p := range netMap.OfflinePeers {
appendByID(p.ID)
accessiblePeers = append(accessiblePeers, peerToAccessiblePeer(p, dnsDomain))
}
return accessiblePeers

View File

@@ -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{}, nil)
networkMapController := controller.NewController(ctx, store, metrics, peersUpdateManager, requestBuffer, server.MockIntegratedValidator{}, settingsManager, "", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peersManager), &config.Config{})
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{}, nil)
networkMapController := controller.NewController(ctx, store, metrics, peersUpdateManager, requestBuffer, server.MockIntegratedValidator{}, settingsManager, "", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peersManager), &config.Config{})
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)

View File

@@ -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{}, nil)
networkMapController := controller.NewController(ctx, testStore, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(testStore, peersManager), &config.Config{})
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

View File

@@ -11,7 +11,6 @@ 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.
@@ -110,7 +109,7 @@ func (am *DefaultAccountManager) GetValidatedPeers(ctx context.Context, accountI
return nil, nil, err
}
validPeers, err := am.integratedPeerValidator.GetValidatedPeers(ctx, accountID, types.TwinGroups(groups), types.TwinPeers(peers), settings.Extra)
validPeers, err := am.integratedPeerValidator.GetValidatedPeers(ctx, accountID, groups, peers, settings.Extra)
if err != nil {
return nil, nil, err
}
@@ -139,7 +138,7 @@ func (a MockIntegratedValidator) ValidatePeer(_ context.Context, update *nbpeer.
return update, false, nil
}
func (a MockIntegratedValidator) GetValidatedPeers(_ context.Context, accountID string, groups []*nmdata.Group, peers []*nmdata.Peer, extraSettings *types.ExtraSettings) (map[string]struct{}, error) {
func (a MockIntegratedValidator) GetValidatedPeers(_ context.Context, accountID string, groups []*types.Group, peers []*nbpeer.Peer, extraSettings *types.ExtraSettings) (map[string]struct{}, error) {
validatedPeers := make(map[string]struct{})
for _, peer := range peers {
validatedPeers[peer.ID] = struct{}{}

View File

@@ -5,7 +5,6 @@ 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"
)
@@ -15,7 +14,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 []*nmdata.Group, peers []*nmdata.Peer, extraSettings *types.ExtraSettings) (map[string]struct{}, error)
GetValidatedPeers(ctx context.Context, accountID string, groups []*types.Group, peers []*nbpeer.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))

View File

@@ -10,7 +10,6 @@ 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"
)
@@ -36,7 +35,7 @@ func (v *IntegratedValidatorImpl) IsNotValidPeer(_ context.Context, _ string, _
return false, false, nil
}
func (v *IntegratedValidatorImpl) GetValidatedPeers(_ context.Context, _ string, _ []*nmdata.Group, peers []*nmdata.Peer, _ *types.ExtraSettings) (map[string]struct{}, error) {
func (v *IntegratedValidatorImpl) GetValidatedPeers(_ context.Context, _ string, _ []*types.Group, peers []*nbpeer.Peer, _ *types.ExtraSettings) (map[string]struct{}, error) {
validatedPeers := make(map[string]struct{})
for _, p := range peers {
validatedPeers[p.ID] = struct{}{}

View File

@@ -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, nil)
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeralMgr, config)
accountManager, err := BuildManager(ctx, nil, store, networkMapController, jobManager, nil, "",
eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)

View File

@@ -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, nil)
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)
accountManager, err := server.BuildManager(
context.Background(),

View File

@@ -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{}, nil)
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{})
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)
}

View File

@@ -21,7 +21,6 @@ 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"
@@ -1589,7 +1588,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 []*nmdata.Peer) {
add := func(peers []*nbpeer.Peer) {
for _, p := range peers {
if p == nil || p.ID == "" || p.ID == selfPeerID {
continue

View File

@@ -57,7 +57,6 @@ 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"
)
@@ -1092,22 +1091,22 @@ func TestToSyncResponse(t *testing.T) {
Signature: "turn-pass",
}
networkMap := &types.NetworkMap{
Network: &nmdata.Network{Net: *ipnet, Serial: 1000},
Peers: []*nmdata.Peer{{
Network: &types.Network{Net: *ipnet, Serial: 1000},
Peers: []*nbpeer.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: []*nmdata.Peer{{
OfflinePeers: []*nbpeer.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: []*nmdata.Route{
Routes: []*nbroute.Route{
{
ID: "route1",
Network: netip.MustParsePrefix("10.0.0.0/24"),
@@ -1181,7 +1180,7 @@ func TestToSyncResponse(t *testing.T) {
}
dnsCache := &cache.DNSConfigCache{}
accountSettings := &types.Settings{RoutingPeerDNSResolutionEnabled: true}
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))
response := grpc.ToSyncResponse(context.Background(), config, config.HttpConfig, config.DeviceAuthorizationFlow, peer, turnRelayToken, turnRelayToken, networkMap, dnsName, checks, dnsCache, accountSettings, nil, []string{}, int64(dnsForwarderPort))
assert.NotNil(t, response)
// assert peer config
@@ -1301,7 +1300,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{}, nil)
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{})
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)
@@ -1392,7 +1391,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{}, nil)
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{})
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)
@@ -1551,7 +1550,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{}, nil)
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{})
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)
@@ -1636,7 +1635,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{}, nil)
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{})
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)

View File

@@ -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, types.TwinRoute(expectedRoute).Equal(peer1Routes.Routes[0]), "received route should be equal")
require.True(t, 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{}, nil)
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{})
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 {

View File

@@ -9,6 +9,7 @@ import (
"strings"
"time"
"github.com/hashicorp/go-multierror"
"github.com/miekg/dns"
"github.com/rs/xid"
log "github.com/sirupsen/logrus"
@@ -17,6 +18,8 @@ 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"
@@ -25,8 +28,6 @@ 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"
)
@@ -282,8 +283,8 @@ func (a *Account) SynthesizePrivateServiceZones(peerID string) []nbdns.CustomZon
// it, adding a single private service would black-hole every
// other name under the zone apex.
zone = &nbdns.CustomZone{
Domain: dns.Fqdn(serviceDomainZone),
Records: []nbdns.SimpleRecord{},
Domain: dns.Fqdn(serviceDomainZone),
Records: []nbdns.SimpleRecord{},
NonAuthoritative: true,
SearchDomainDisabled: true,
}
@@ -381,11 +382,94 @@ func peerInDistributionGroups(peerGroups LookupMap, distributionGroups []string)
}
func (a *Account) GetPeersCustomZone(ctx context.Context, dnsDomain string) nbdns.CustomZone {
twins := make(map[string]*nmdata.Peer, len(a.Peers))
for id, p := range a.Peers {
twins[id] = twinPeer(p)
var merr *multierror.Error
if dnsDomain == "" {
log.WithContext(ctx).Error("no dns domain is set, returning empty zone")
return nbdns.CustomZone{}
}
return fromTwinCustomZone(networkmap.PeersCustomZone(ctx, a.Id, dnsDomain, twins, a.peerIPv6AllowedSet()))
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
}
// GetExpiredPeers returns peers that have been expired
@@ -978,54 +1062,6 @@ 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 {
@@ -1762,3 +1798,66 @@ 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
}

View File

@@ -2,14 +2,18 @@ package types
import (
"context"
"slices"
"time"
log "github.com/sirupsen/logrus"
nbdns "github.com/netbirdio/netbird/dns"
"github.com/netbirdio/netbird/management/internals/modules/zones"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/telemetry"
"github.com/netbirdio/netbird/route"
)
// GetPeerNetworkMapResult dispatches to either the legacy-NetworkMap path or
@@ -90,9 +94,6 @@ 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,
@@ -103,10 +104,617 @@ func (a *Account) GetPeerNetworkMapComponents(
routers map[string]map[string]*routerTypes.NetworkRouter,
groupIDToUserIDs map[string][]string,
) *NetworkMapComponents {
nmd := a.toNetworkMapData(accountZones, validatedPeersMap, resourcePolicies, routers, groupIDToUserIDs)
components := nmd.GetPeerNetworkMapComponents(peerID, TwinCustomZone(peersCustomZone))
if components != nil {
components.ForceRoutingPeerDNSResolution = a.forcesRoutingPeerDNSResolution(peerID, routers)
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]*nbpeer.Peer{peerID: 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]*nbpeer.Peer{peerID: 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]*routerTypes.NetworkRouter),
NetworkResources: make([]*resourceTypes.NetworkResource, 0),
PostureFailedPeers: make(map[string]map[string]struct{}, len(a.PostureChecks)),
RouterPeers: make(map[string]*nbpeer.Peer),
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 = 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)
}
}
} 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
}
}
}
for _, dstGroupID := range rule.Destinations {
if g := a.Groups[dstGroupID]; g != nil {
if _, exists := components.Groups[dstGroupID]; !exists {
components.Groups[dstGroupID] = 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] = networkRoutingPeers
for peerIDKey := range networkRoutingPeers {
if p := a.Peers[peerIDKey]; p != nil {
if _, exists := components.RouterPeers[peerIDKey]; !exists {
components.RouterPeers[peerIDKey] = p
}
if _, exists := components.Peers[peerIDKey]; !exists {
if _, validated := validatedPeersMap[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
}
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]*nbpeer.Peer, map[string]*Group, []*Policy, []*route.Route, sshRequirements) {
relevantPeerIDs := make(map[string]*nbpeer.Peer, 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)
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
}
}
}
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
}
}
}
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 {
relevantPeerIDs[pid] = a.GetPeer(pid)
}
for _, dstGroupID := range rule.Destinations {
relevantGroupIDs[dstGroupID] = a.GetGroup(dstGroupID)
}
}
if peerInDestinations {
policyRelevant = true
for _, pid := range sourcePeers {
relevantPeerIDs[pid] = 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 (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]*Group, peers map[string]*nbpeer.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 []*Policy, resourcePoliciesMap map[string][]*Policy, peers map[string]*nbpeer.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 []nbdns.SimpleRecord, peers map[string]*nbpeer.Peer, 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
}

View File

@@ -1,566 +0,0 @@
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,
}
}

View File

@@ -9,7 +9,7 @@ import (
"github.com/stretchr/testify/require"
nbdns "github.com/netbirdio/netbird/dns"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
)
func TestPrivateService_NetworkMap_UserPeer_AndProxyPeer(t *testing.T) {
@@ -49,7 +49,7 @@ func TestPrivateService_NetworkMap_UserPeer_AndProxyPeer(t *testing.T) {
})
}
func netmapPeerIDs(peers []*nmdata.Peer) []string {
func netmapPeerIDs(peers []*nbpeer.Peer) []string {
ids := make([]string, 0, len(peers))
for _, p := range peers {
if p == nil {

View File

@@ -13,6 +13,8 @@ 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"
@@ -1038,6 +1040,518 @@ 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()

View File

@@ -67,16 +67,12 @@ func PolicyRuleImpliesLegacySSH(rule *PolicyRule) bool {
return sharedtypes.PolicyRuleImpliesLegacySSH(rule)
}
// 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))
return sharedtypes.ExpandPortsAndRanges(base, rule, peer)
}
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)
return sharedtypes.AppendIPv6FirewallRule(rules, rulesExists, peer, targetPeer, rule, rc)
}
func CalculateNetworkMapFromComponents(ctx context.Context, components *NetworkMapComponents) *NetworkMap {
@@ -84,7 +80,7 @@ func CalculateNetworkMapFromComponents(ctx context.Context, components *NetworkM
}
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)
return sharedtypes.GenerateRouteFirewallRules(ctx, route, rule, groupPeers, direction, includeIPv6)
}
func AllocateIPv6Subnet(r *rand.Rand) net.IPNet {

View File

@@ -9,7 +9,6 @@ import (
"github.com/stretchr/testify/require"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
func TestNetworkMapComponents_IPv6EndToEnd(t *testing.T) {
@@ -105,7 +104,7 @@ func TestNetworkMapComponents_RemotePeerWithoutCapability(t *testing.T) {
require.NotNil(t, nm)
t.Run("AllowedIPs include remote v6", func(t *testing.T) {
var dst *nmdata.Peer
var dst *nbpeer.Peer
for _, p := range nm.Peers {
if p.ID == "peer-dst-1" {
dst = p

View File

@@ -1,703 +0,0 @@
//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
}

View File

@@ -1,48 +0,0 @@
//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
)

View File

@@ -1,105 +0,0 @@
//go:build nmapequiv
package legacynmap
import (
"net/netip"
"time"
)
// ComponentPeer is the self-contained peer representation used by
// NetworkMapComponents and the calculated NetworkMap. It carries exactly the
// subset of peer data that crosses the components wire format, so the shared
// calculation layer stays independent of the management server's domain
// types.
type ComponentPeer struct {
ID string
Key string
IP netip.Addr
IPv6 netip.Addr
DNSLabel string
SSHKey string
SSHEnabled bool
ServerSSHAllowed bool
AgentVersion string
SupportsSourcePrefixes bool
SupportsIPv6 bool
LoginExpirationEnabled bool
AddedWithSSOLogin bool
LastLogin time.Time
}
// FQDN returns the peer's FQDN combined of the peer's DNS label and the system's DNS domain.
func (p *ComponentPeer) FQDN(dnsDomain string) string {
if dnsDomain == "" {
return ""
}
return p.DNSLabel + "." + dnsDomain
}
// LoginExpired indicates whether the peer's login has expired, mirroring the
// server-side peer semantics: only SSO-added peers with login expiration
// enabled can expire.
func (p *ComponentPeer) LoginExpired(expiresIn time.Duration) (bool, time.Duration) {
if !p.AddedWithSSOLogin || !p.LoginExpirationEnabled {
return false, 0
}
timeLeft := time.Until(p.LastLogin.Add(expiresIn))
return timeLeft <= 0, timeLeft
}
// GroupAllName is the reserved name of the default group that contains every peer in an account.
const GroupAllName = "All"
// ComponentGroup is the self-contained group representation used by
// NetworkMapComponents: just the membership view the network-map calculation
// needs, without the server's storage fields.
type ComponentGroup struct {
ID string
PublicID string
Name string
Peers []string
}
// IsGroupAll checks if the group is a default "All" group.
func (g *ComponentGroup) IsGroupAll() bool {
return g.Name == GroupAllName
}
// ComponentRouter is the self-contained network-router representation used by
// NetworkMapComponents.
type ComponentRouter struct {
NetworkID string
PublicID string
Peer string
PeerGroups []string
Masquerade bool
Metric int
Enabled bool
}
// ComponentResourceType mirrors the network-resource type enum on the
// components wire format.
type ComponentResourceType string
const (
ComponentResourceHost ComponentResourceType = "host"
ComponentResourceSubnet ComponentResourceType = "subnet"
ComponentResourceDomain ComponentResourceType = "domain"
)
// ComponentResource is the self-contained network-resource representation
// used by NetworkMapComponents.
type ComponentResource struct {
ID string
PublicID string
NetworkID string
AccountID string
Name string
Description string
Type ComponentResourceType
Address string
Domain string
Prefix netip.Prefix
Enabled bool
}

View File

@@ -1,128 +0,0 @@
//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,
}
}

View File

@@ -1,284 +0,0 @@
//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
}

View File

@@ -1,7 +0,0 @@
// 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

View File

@@ -1,570 +0,0 @@
//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
}

View File

@@ -1,157 +0,0 @@
//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]
}

File diff suppressed because it is too large Load Diff

View File

@@ -1,208 +0,0 @@
//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)
}

View File

@@ -18,7 +18,6 @@ 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 {
@@ -50,7 +49,7 @@ func allPeersValidated(account *types.Account, excludePeerIDs ...string) map[str
return validated
}
func peerIDs(peers []*nmdata.Peer) []string {
func peerIDs(peers []*nbpeer.Peer) []string {
ids := make([]string, len(peers))
for i, p := range peers {
ids[i] = p.ID
@@ -626,7 +625,7 @@ func TestNetworkMapComponents_DomainNetworkResource(t *testing.T) {
var hasDomainRoute bool
for _, r := range nm.Routes {
if r.NetworkType == int(route.DomainNetwork) && len(r.Domains) > 0 && r.Domains[0].SafeString() == "api.example.com" {
if r.NetworkType == route.DomainNetwork && len(r.Domains) > 0 && r.Domains[0].SafeString() == "api.example.com" {
hasDomainRoute = true
}
}

View File

@@ -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, types.TwinPeer(peer), nil, nil, networkMap, "netbird.cloud", nil, dnsCache, types.TwinAccountSettings(settings), nil, nil, 0)
legacyResp := mgmtgrpc.ToSyncResponse(ctx, nil, nil, nil, peer, nil, nil, networkMap, "netbird.cloud", nil, dnsCache, 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, types.TwinPeer(peer), nil, nil, networkMap, "netbird.cloud", nil, dnsCache, types.TwinAccountSettings(settings), nil, nil, 0)
resp := mgmtgrpc.ToSyncResponse(ctx, nil, nil, nil, peer, nil, nil, networkMap, "netbird.cloud", nil, dnsCache, 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, types.TwinPeer(peer), nil, nil, networkMap, "netbird.cloud", nil, dnsCache, types.TwinAccountSettings(settings), nil, nil, 0)
legacyResp := mgmtgrpc.ToSyncResponse(ctx, nil, nil, nil, peer, nil, nil, networkMap, "netbird.cloud", nil, dnsCache, settings, nil, nil, 0)
legacyBytes, err := goproto.Marshal(legacyResp.NetworkMap)
if err != nil {
b.Fatalf("marshal legacy networkmap: %v", err)

View File

@@ -45,7 +45,7 @@ func TestNetworkMapWireBreakdown(t *testing.T) {
dnsCache := &cache.DNSConfigCache{}
settings := &types.Settings{}
legacyResp := mgmtgrpc.ToSyncResponse(ctx, nil, nil, nil, types.TwinPeer(peer), nil, nil, networkMap, "netbird.cloud", nil, dnsCache, types.TwinAccountSettings(settings), nil, nil, 0)
legacyResp := mgmtgrpc.ToSyncResponse(ctx, nil, nil, nil, peer, nil, nil, networkMap, "netbird.cloud", nil, dnsCache, settings, nil, nil, 0)
legacyTotal := mustMarshalSize(t, legacyResp.NetworkMap)
envelope := mgmtgrpc.EncodeNetworkMapEnvelope(mgmtgrpc.ComponentsEnvelopeInput{

View File

@@ -10,7 +10,6 @@ COPY encryption ./encryption
COPY flow ./flow
COPY formatter ./formatter
COPY monotime ./monotime
COPY management ./management
COPY proxy ./proxy
COPY route ./route
COPY shared ./shared

View File

@@ -161,16 +161,6 @@ openai:
input_per_1k: 0.0001
output_per_1k: 0
# Kimi / Moonshot AI (kimi_api) — OpenAI-compatible /v1 endpoint. Moonshot
# reports cache hits OpenAI-style when present; cached input is 10% of
# input ($0.30 vs $3.00 per MTok). kimi-k3 is the only model the platform
# serves newer accounts (K2-era ids and kimi-latest 404), matching the
# management catalog.
kimi-k3:
input_per_1k: 0.003
output_per_1k: 0.015
cached_input_per_1k: 0.0003
anthropic:
# Claude 4.x family — cache reads ≈10% of input, cache writes ≈125% of input.
# Pricing source: Anthropic's current published rates per million tokens,
@@ -216,20 +206,6 @@ anthropic:
cache_read_per_1k: 0.0001
cache_creation_per_1k: 0.00125
# Kimi / Moonshot AI (kimi_api) via the Anthropic-compatible endpoint
# (/anthropic/v1/messages — the official Claude Code setup). Same rates
# as the OpenAI-shape entry above. "kimi-k3[1m]" is the model id some
# Claude Code guides set for the 1M-context alias; priced identically so
# cost metering doesn't silently skip those requests.
kimi-k3:
input_per_1k: 0.003
output_per_1k: 0.015
cache_read_per_1k: 0.0003
"kimi-k3[1m]":
input_per_1k: 0.003
output_per_1k: 0.015
cache_read_per_1k: 0.0003
bedrock:
# AWS Bedrock model ids, normalised by the request parser (cross-region
# inference-profile prefix + version/throughput suffix stripped), e.g.

View File

@@ -3,8 +3,8 @@
# FreeBSD Port Diff Generator for NetBird
#
# This script generates the diff file required for submitting a FreeBSD port update.
# It works on macOS, Linux, and FreeBSD by fetching files from the FreeBSD ports
# GitHub mirror and computing checksums from the Go module proxy.
# It works on macOS, Linux, and FreeBSD by fetching files from FreeBSD cgit and
# computing checksums from the Go module proxy.
#
# Usage: ./freebsd-port-diff.sh [new_version]
# Example: ./freebsd-port-diff.sh 0.60.7
@@ -14,7 +14,7 @@
set -e
GITHUB_REPO="netbirdio/netbird"
PORTS_MIRROR_BASE="https://raw.githubusercontent.com/freebsd/freebsd-ports/main/security/netbird"
PORTS_CGIT_BASE="https://cgit.freebsd.org/ports/plain/security/netbird"
GO_PROXY="https://proxy.golang.org/github.com/netbirdio/netbird/@v"
OUTPUT_DIR="${OUTPUT_DIR:-.}"
AWK_FIRST_FIELD='{print $1}'
@@ -30,17 +30,10 @@ fetch_all_tags() {
fetch_current_ports_version() {
echo "Fetching current version from FreeBSD ports..." >&2
local makefile version
makefile=$(fetch_ports_file "Makefile") || return 1
version=$(echo "$makefile" | \
curl -sL "${PORTS_CGIT_BASE}/Makefile" 2>/dev/null | \
grep -E "^DISTVERSION=" | \
sed 's/DISTVERSION=[[:space:]]*//' | \
tr -d '\t ')
if [[ -z "$version" ]]; then
echo "Error: Could not extract DISTVERSION from ports Makefile" >&2
return 1
fi
echo "$version"
tr -d '\t '
return 0
}
@@ -52,16 +45,7 @@ fetch_latest_github_release() {
fetch_ports_file() {
local filename="$1"
local content
if ! content=$(curl -fsL --proto '=https' --proto-redir '=https' --retry 3 "${PORTS_MIRROR_BASE}/${filename}" 2>/dev/null); then
echo "Error: Could not fetch ${filename} from ${PORTS_MIRROR_BASE}" >&2
return 1
fi
if [[ "$content" == \<* ]]; then
echo "Error: Received HTML instead of ${filename} from ${PORTS_MIRROR_BASE}" >&2
return 1
fi
printf '%s' "$content"
curl -sL "${PORTS_CGIT_BASE}/${filename}" 2>/dev/null
return 0
}

Some files were not shown because too many files have changed in this diff Show More