diff --git a/client/cmd/debug.go b/client/cmd/debug.go
index 0e2717756..e3d3afe5f 100644
--- a/client/cmd/debug.go
+++ b/client/cmd/debug.go
@@ -199,9 +199,11 @@ func runForDuration(cmd *cobra.Command, args []string) error {
cmd.Println("Log level set to trace.")
}
+ needsRestoreUp := false
if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil {
cmd.PrintErrf("Failed to bring service down: %v\n", status.Convert(err).Message())
} else {
+ needsRestoreUp = !stateWasDown
cmd.Println("netbird down")
}
@@ -217,6 +219,7 @@ func runForDuration(cmd *cobra.Command, args []string) error {
if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil {
cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message())
} else {
+ needsRestoreUp = false
cmd.Println("netbird up")
}
@@ -264,6 +267,14 @@ func runForDuration(cmd *cobra.Command, args []string) error {
return fmt.Errorf("failed to bundle debug: %v", status.Convert(err).Message())
}
+ if needsRestoreUp {
+ if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil {
+ cmd.PrintErrf("Failed to restore service up state: %v\n", status.Convert(err).Message())
+ } else {
+ cmd.Println("netbird up (restored)")
+ }
+ }
+
if stateWasDown {
if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil {
cmd.PrintErrf("Failed to restore service down state: %v\n", status.Convert(err).Message())
diff --git a/client/firewall/create_linux.go b/client/firewall/create_linux.go
index 12dcaee8a..d781ebd77 100644
--- a/client/firewall/create_linux.go
+++ b/client/firewall/create_linux.go
@@ -6,6 +6,7 @@ import (
"errors"
"fmt"
"os"
+ "strconv"
"github.com/coreos/go-iptables/iptables"
"github.com/google/nftables"
@@ -35,20 +36,27 @@ const SKIP_NFTABLES_ENV = "NB_SKIP_NFTABLES_CHECK"
type FWType int
func NewFirewall(iface IFaceMapper, stateManager *statemanager.Manager, flowLogger nftypes.FlowLogger, disableServerRoutes bool, mtu uint16) (firewall.Manager, error) {
- // on the linux system we try to user nftables or iptables
- // in any case, because we need to allow netbird interface traffic
- // so we use AllowNetbird traffic from these firewall managers
- // for the userspace packet filtering firewall
+ // We run in userspace mode and force userspace firewall was requested. We don't attempt native firewall.
+ if iface.IsUserspaceBind() && forceUserspaceFirewall() {
+ log.Info("forcing userspace firewall")
+ return createUserspaceFirewall(iface, nil, disableServerRoutes, flowLogger, mtu)
+ }
+
+ // Use native firewall for either kernel or userspace, the interface appears identical to netfilter
fm, err := createNativeFirewall(iface, stateManager, disableServerRoutes, mtu)
+ // Kernel cannot fall back to anything else, need to return error
if !iface.IsUserspaceBind() {
return fm, err
}
+ // Fall back to the userspace packet filter if native is unavailable
if err != nil {
log.Warnf("failed to create native firewall: %v. Proceeding with userspace", err)
+ return createUserspaceFirewall(iface, nil, disableServerRoutes, flowLogger, mtu)
}
- return createUserspaceFirewall(iface, fm, disableServerRoutes, flowLogger, mtu)
+
+ return fm, nil
}
func createNativeFirewall(iface IFaceMapper, stateManager *statemanager.Manager, routes bool, mtu uint16) (firewall.Manager, error) {
@@ -160,3 +168,17 @@ func isIptablesClientAvailable(client *iptables.IPTables) bool {
_, err := client.ListChains("filter")
return err == nil
}
+
+func forceUserspaceFirewall() bool {
+ val := os.Getenv(EnvForceUserspaceFirewall)
+ if val == "" {
+ return false
+ }
+
+ force, err := strconv.ParseBool(val)
+ if err != nil {
+ log.Warnf("failed to parse %s: %v", EnvForceUserspaceFirewall, err)
+ return false
+ }
+ return force
+}
diff --git a/client/firewall/iface.go b/client/firewall/iface.go
index b83c5f912..491f03269 100644
--- a/client/firewall/iface.go
+++ b/client/firewall/iface.go
@@ -7,6 +7,12 @@ import (
"github.com/netbirdio/netbird/client/iface/wgaddr"
)
+// EnvForceUserspaceFirewall forces the use of the userspace packet filter even when
+// native iptables/nftables is available. This only applies when the WireGuard interface
+// runs in userspace mode. When set, peer ACLs are handled by USPFilter instead of
+// kernel netfilter rules.
+const EnvForceUserspaceFirewall = "NB_FORCE_USERSPACE_FIREWALL"
+
// IFaceMapper defines subset methods of interface required for manager
type IFaceMapper interface {
Name() string
diff --git a/client/firewall/iptables/manager_linux.go b/client/firewall/iptables/manager_linux.go
index 2fc6f8ec8..a1d4467d5 100644
--- a/client/firewall/iptables/manager_linux.go
+++ b/client/firewall/iptables/manager_linux.go
@@ -33,7 +33,6 @@ type Manager struct {
type iFaceMapper interface {
Name() string
Address() wgaddr.Address
- IsUserspaceBind() bool
}
// Create iptables firewall manager
@@ -64,10 +63,9 @@ func Create(wgIface iFaceMapper, mtu uint16) (*Manager, error) {
func (m *Manager) Init(stateManager *statemanager.Manager) error {
state := &ShutdownState{
InterfaceState: &InterfaceState{
- NameStr: m.wgIface.Name(),
- WGAddress: m.wgIface.Address(),
- UserspaceBind: m.wgIface.IsUserspaceBind(),
- MTU: m.router.mtu,
+ NameStr: m.wgIface.Name(),
+ WGAddress: m.wgIface.Address(),
+ MTU: m.router.mtu,
},
}
stateManager.RegisterState(state)
@@ -203,12 +201,10 @@ func (m *Manager) Close(stateManager *statemanager.Manager) error {
return nberrors.FormatErrorOrNil(merr)
}
-// AllowNetbird allows netbird interface traffic
+// AllowNetbird allows netbird interface traffic.
+// This is called when USPFilter wraps the native firewall, adding blanket accept
+// rules so that packet filtering is handled in userspace instead of by netfilter.
func (m *Manager) AllowNetbird() error {
- if !m.wgIface.IsUserspaceBind() {
- return nil
- }
-
_, err := m.AddPeerFiltering(
nil,
net.IP{0, 0, 0, 0},
diff --git a/client/firewall/iptables/manager_linux_test.go b/client/firewall/iptables/manager_linux_test.go
index ee47a27c0..cc4bda0e0 100644
--- a/client/firewall/iptables/manager_linux_test.go
+++ b/client/firewall/iptables/manager_linux_test.go
@@ -47,8 +47,6 @@ func (i *iFaceMock) Address() wgaddr.Address {
panic("AddressFunc is not set")
}
-func (i *iFaceMock) IsUserspaceBind() bool { return false }
-
func TestIptablesManager(t *testing.T) {
ipv4Client, err := iptables.NewWithProtocol(iptables.ProtocolIPv4)
require.NoError(t, err)
diff --git a/client/firewall/iptables/state_linux.go b/client/firewall/iptables/state_linux.go
index c88774c1f..121c755e9 100644
--- a/client/firewall/iptables/state_linux.go
+++ b/client/firewall/iptables/state_linux.go
@@ -9,10 +9,9 @@ import (
)
type InterfaceState struct {
- NameStr string `json:"name"`
- WGAddress wgaddr.Address `json:"wg_address"`
- UserspaceBind bool `json:"userspace_bind"`
- MTU uint16 `json:"mtu"`
+ NameStr string `json:"name"`
+ WGAddress wgaddr.Address `json:"wg_address"`
+ MTU uint16 `json:"mtu"`
}
func (i *InterfaceState) Name() string {
@@ -23,10 +22,6 @@ func (i *InterfaceState) Address() wgaddr.Address {
return i.WGAddress
}
-func (i *InterfaceState) IsUserspaceBind() bool {
- return i.UserspaceBind
-}
-
type ShutdownState struct {
sync.Mutex
diff --git a/client/firewall/nftables/manager_linux.go b/client/firewall/nftables/manager_linux.go
index beb5b70a7..0b5b61e04 100644
--- a/client/firewall/nftables/manager_linux.go
+++ b/client/firewall/nftables/manager_linux.go
@@ -40,7 +40,6 @@ func getTableName() string {
type iFaceMapper interface {
Name() string
Address() wgaddr.Address
- IsUserspaceBind() bool
}
// Manager of iptables firewall
@@ -106,10 +105,9 @@ func (m *Manager) Init(stateManager *statemanager.Manager) error {
// cleanup using Close() without needing to store specific rules.
if err := stateManager.UpdateState(&ShutdownState{
InterfaceState: &InterfaceState{
- NameStr: m.wgIface.Name(),
- WGAddress: m.wgIface.Address(),
- UserspaceBind: m.wgIface.IsUserspaceBind(),
- MTU: m.router.mtu,
+ NameStr: m.wgIface.Name(),
+ WGAddress: m.wgIface.Address(),
+ MTU: m.router.mtu,
},
}); err != nil {
log.Errorf("failed to update state: %v", err)
@@ -205,12 +203,10 @@ func (m *Manager) RemoveNatRule(pair firewall.RouterPair) error {
return m.router.RemoveNatRule(pair)
}
-// AllowNetbird allows netbird interface traffic
+// AllowNetbird allows netbird interface traffic.
+// This is called when USPFilter wraps the native firewall, adding blanket accept
+// rules so that packet filtering is handled in userspace instead of by netfilter.
func (m *Manager) AllowNetbird() error {
- if !m.wgIface.IsUserspaceBind() {
- return nil
- }
-
m.mutex.Lock()
defer m.mutex.Unlock()
diff --git a/client/firewall/nftables/manager_linux_test.go b/client/firewall/nftables/manager_linux_test.go
index 75b1e2b6c..d48e4ba88 100644
--- a/client/firewall/nftables/manager_linux_test.go
+++ b/client/firewall/nftables/manager_linux_test.go
@@ -52,8 +52,6 @@ func (i *iFaceMock) Address() wgaddr.Address {
panic("AddressFunc is not set")
}
-func (i *iFaceMock) IsUserspaceBind() bool { return false }
-
func TestNftablesManager(t *testing.T) {
// just check on the local interface
diff --git a/client/firewall/nftables/state_linux.go b/client/firewall/nftables/state_linux.go
index 48b7b3741..462ad2556 100644
--- a/client/firewall/nftables/state_linux.go
+++ b/client/firewall/nftables/state_linux.go
@@ -8,10 +8,9 @@ import (
)
type InterfaceState struct {
- NameStr string `json:"name"`
- WGAddress wgaddr.Address `json:"wg_address"`
- UserspaceBind bool `json:"userspace_bind"`
- MTU uint16 `json:"mtu"`
+ NameStr string `json:"name"`
+ WGAddress wgaddr.Address `json:"wg_address"`
+ MTU uint16 `json:"mtu"`
}
func (i *InterfaceState) Name() string {
@@ -22,10 +21,6 @@ func (i *InterfaceState) Address() wgaddr.Address {
return i.WGAddress
}
-func (i *InterfaceState) IsUserspaceBind() bool {
- return i.UserspaceBind
-}
-
type ShutdownState struct {
InterfaceState *InterfaceState `json:"interface_state,omitempty"`
}
diff --git a/client/internal/acl/manager_test.go b/client/internal/acl/manager_test.go
index bd7adfaef..408ed992f 100644
--- a/client/internal/acl/manager_test.go
+++ b/client/internal/acl/manager_test.go
@@ -19,6 +19,9 @@ import (
var flowLogger = netflow.NewManager(nil, []byte{}, nil).GetLogger()
func TestDefaultManager(t *testing.T) {
+ t.Setenv("NB_WG_KERNEL_DISABLED", "true")
+ t.Setenv(firewall.EnvForceUserspaceFirewall, "true")
+
networkMap := &mgmProto.NetworkMap{
FirewallRules: []*mgmProto.FirewallRule{
{
@@ -135,6 +138,7 @@ func TestDefaultManager(t *testing.T) {
func TestDefaultManagerStateless(t *testing.T) {
// stateless currently only in userspace, so we have to disable kernel
t.Setenv("NB_WG_KERNEL_DISABLED", "true")
+ t.Setenv(firewall.EnvForceUserspaceFirewall, "true")
t.Setenv("NB_DISABLE_CONNTRACK", "true")
networkMap := &mgmProto.NetworkMap{
@@ -194,6 +198,7 @@ func TestDefaultManagerStateless(t *testing.T) {
// This tests the full ACL manager -> uspfilter integration.
func TestDenyRulesNotAccumulatedOnRepeatedApply(t *testing.T) {
t.Setenv("NB_WG_KERNEL_DISABLED", "true")
+ t.Setenv(firewall.EnvForceUserspaceFirewall, "true")
networkMap := &mgmProto.NetworkMap{
FirewallRules: []*mgmProto.FirewallRule{
@@ -258,6 +263,7 @@ func TestDenyRulesNotAccumulatedOnRepeatedApply(t *testing.T) {
// up when they're removed from the network map in a subsequent update.
func TestDenyRulesCleanedUpOnRemoval(t *testing.T) {
t.Setenv("NB_WG_KERNEL_DISABLED", "true")
+ t.Setenv(firewall.EnvForceUserspaceFirewall, "true")
ctrl := gomock.NewController(t)
defer ctrl.Finish()
@@ -339,6 +345,7 @@ func TestDenyRulesCleanedUpOnRemoval(t *testing.T) {
// one added without leaking.
func TestRuleUpdateChangingAction(t *testing.T) {
t.Setenv("NB_WG_KERNEL_DISABLED", "true")
+ t.Setenv(firewall.EnvForceUserspaceFirewall, "true")
ctrl := gomock.NewController(t)
defer ctrl.Finish()
diff --git a/client/internal/debug/debug.go b/client/internal/debug/debug.go
index c9ebf25e5..6a8eae324 100644
--- a/client/internal/debug/debug.go
+++ b/client/internal/debug/debug.go
@@ -25,6 +25,7 @@ import (
"google.golang.org/protobuf/encoding/protojson"
"github.com/netbirdio/netbird/client/anonymize"
+ "github.com/netbirdio/netbird/client/configs"
"github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/internal/updater/installer"
@@ -52,6 +53,7 @@ resolved_domains.txt: Anonymized resolved domain IP addresses from the status re
config.txt: Anonymized configuration information of the NetBird client.
network_map.json: Anonymized sync response containing peer configurations, routes, DNS settings, and firewall rules.
state.json: Anonymized client state dump containing netbird states for the active profile.
+service_params.json: Sanitized service install parameters (service.json). Sensitive environment variable values are masked. Only present when service.json exists.
metrics.txt: Buffered client metrics in InfluxDB line protocol format. Only present when metrics collection is enabled. Peer identifiers are anonymized.
mutex.prof: Mutex profiling information.
goroutine.prof: Goroutine profiling information.
@@ -359,6 +361,10 @@ func (g *BundleGenerator) createArchive() error {
log.Errorf("failed to add corrupted state files to debug bundle: %v", err)
}
+ if err := g.addServiceParams(); err != nil {
+ log.Errorf("failed to add service params to debug bundle: %v", err)
+ }
+
if err := g.addMetrics(); err != nil {
log.Errorf("failed to add metrics to debug bundle: %v", err)
}
@@ -488,6 +494,90 @@ func (g *BundleGenerator) addConfig() error {
return nil
}
+const (
+ serviceParamsFile = "service.json"
+ serviceParamsBundle = "service_params.json"
+ maskedValue = "***"
+ envVarPrefix = "NB_"
+ jsonKeyManagementURL = "management_url"
+ jsonKeyServiceEnv = "service_env_vars"
+)
+
+var sensitiveEnvSubstrings = []string{"key", "token", "secret", "password", "credential"}
+
+// addServiceParams reads the service.json file and adds a sanitized version to the bundle.
+// Non-NB_ env vars and vars with sensitive names are masked. Other NB_ values are anonymized.
+func (g *BundleGenerator) addServiceParams() error {
+ path := filepath.Join(configs.StateDir, serviceParamsFile)
+
+ data, err := os.ReadFile(path)
+ if err != nil {
+ if os.IsNotExist(err) {
+ return nil
+ }
+ return fmt.Errorf("read service params: %w", err)
+ }
+
+ var params map[string]any
+ if err := json.Unmarshal(data, ¶ms); err != nil {
+ return fmt.Errorf("parse service params: %w", err)
+ }
+
+ if g.anonymize {
+ if mgmtURL, ok := params[jsonKeyManagementURL].(string); ok && mgmtURL != "" {
+ params[jsonKeyManagementURL] = g.anonymizer.AnonymizeURI(mgmtURL)
+ }
+ }
+
+ g.sanitizeServiceEnvVars(params)
+
+ sanitizedData, err := json.MarshalIndent(params, "", " ")
+ if err != nil {
+ return fmt.Errorf("marshal sanitized service params: %w", err)
+ }
+
+ if err := g.addFileToZip(bytes.NewReader(sanitizedData), serviceParamsBundle); err != nil {
+ return fmt.Errorf("add service params to zip: %w", err)
+ }
+
+ return nil
+}
+
+// sanitizeServiceEnvVars masks or anonymizes env var values in service params.
+// Non-NB_ vars and vars with sensitive names (key, token, etc.) are fully masked.
+// Other NB_ var values are passed through the anonymizer when anonymization is enabled.
+func (g *BundleGenerator) sanitizeServiceEnvVars(params map[string]any) {
+ envVars, ok := params[jsonKeyServiceEnv].(map[string]any)
+ if !ok {
+ return
+ }
+
+ sanitized := make(map[string]any, len(envVars))
+ for k, v := range envVars {
+ val, _ := v.(string)
+ switch {
+ case !strings.HasPrefix(k, envVarPrefix) || isSensitiveEnvVar(k):
+ sanitized[k] = maskedValue
+ case g.anonymize:
+ sanitized[k] = g.anonymizer.AnonymizeString(val)
+ default:
+ sanitized[k] = val
+ }
+ }
+ params[jsonKeyServiceEnv] = sanitized
+}
+
+// isSensitiveEnvVar returns true for env var names that may contain secrets.
+func isSensitiveEnvVar(key string) bool {
+ lower := strings.ToLower(key)
+ for _, s := range sensitiveEnvSubstrings {
+ if strings.Contains(lower, s) {
+ return true
+ }
+ }
+ return false
+}
+
func (g *BundleGenerator) addCommonConfigFields(configContent *strings.Builder) {
configContent.WriteString("NetBird Client Configuration:\n\n")
diff --git a/client/internal/debug/debug_test.go b/client/internal/debug/debug_test.go
index 59837c328..6b5bb911c 100644
--- a/client/internal/debug/debug_test.go
+++ b/client/internal/debug/debug_test.go
@@ -1,8 +1,12 @@
package debug
import (
+ "archive/zip"
+ "bytes"
"encoding/json"
"net"
+ "os"
+ "path/filepath"
"strings"
"testing"
@@ -10,6 +14,7 @@ import (
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/anonymize"
+ "github.com/netbirdio/netbird/client/configs"
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
)
@@ -420,6 +425,226 @@ func TestAnonymizeNetworkMap(t *testing.T) {
}
}
+func TestIsSensitiveEnvVar(t *testing.T) {
+ tests := []struct {
+ key string
+ sensitive bool
+ }{
+ {"NB_SETUP_KEY", true},
+ {"NB_API_TOKEN", true},
+ {"NB_CLIENT_SECRET", true},
+ {"NB_PASSWORD", true},
+ {"NB_CREDENTIAL", true},
+ {"NB_LOG_LEVEL", false},
+ {"NB_MANAGEMENT_URL", false},
+ {"NB_HOSTNAME", false},
+ {"HOME", false},
+ {"PATH", false},
+ }
+ for _, tt := range tests {
+ t.Run(tt.key, func(t *testing.T) {
+ assert.Equal(t, tt.sensitive, isSensitiveEnvVar(tt.key))
+ })
+ }
+}
+
+func TestSanitizeServiceEnvVars(t *testing.T) {
+ tests := []struct {
+ name string
+ anonymize bool
+ input map[string]any
+ check func(t *testing.T, params map[string]any)
+ }{
+ {
+ name: "no env vars key",
+ anonymize: false,
+ input: map[string]any{"management_url": "https://mgmt.example.com"},
+ check: func(t *testing.T, params map[string]any) {
+ t.Helper()
+ assert.Equal(t, "https://mgmt.example.com", params["management_url"], "non-env fields should be untouched")
+ _, ok := params[jsonKeyServiceEnv]
+ assert.False(t, ok, "service_env_vars should not be added")
+ },
+ },
+ {
+ name: "non-NB vars are masked",
+ anonymize: false,
+ input: map[string]any{
+ jsonKeyServiceEnv: map[string]any{
+ "HOME": "/root",
+ "PATH": "/usr/bin",
+ "NB_LOG_LEVEL": "debug",
+ },
+ },
+ check: func(t *testing.T, params map[string]any) {
+ t.Helper()
+ env := params[jsonKeyServiceEnv].(map[string]any)
+ assert.Equal(t, maskedValue, env["HOME"], "non-NB_ var should be masked")
+ assert.Equal(t, maskedValue, env["PATH"], "non-NB_ var should be masked")
+ assert.Equal(t, "debug", env["NB_LOG_LEVEL"], "safe NB_ var should pass through")
+ },
+ },
+ {
+ name: "sensitive NB vars are masked",
+ anonymize: false,
+ input: map[string]any{
+ jsonKeyServiceEnv: map[string]any{
+ "NB_SETUP_KEY": "abc123",
+ "NB_API_TOKEN": "tok_xyz",
+ "NB_LOG_LEVEL": "info",
+ },
+ },
+ check: func(t *testing.T, params map[string]any) {
+ t.Helper()
+ env := params[jsonKeyServiceEnv].(map[string]any)
+ assert.Equal(t, maskedValue, env["NB_SETUP_KEY"], "sensitive NB_ var should be masked")
+ assert.Equal(t, maskedValue, env["NB_API_TOKEN"], "sensitive NB_ var should be masked")
+ assert.Equal(t, "info", env["NB_LOG_LEVEL"], "safe NB_ var should pass through")
+ },
+ },
+ {
+ name: "safe NB vars anonymized when anonymize is true",
+ anonymize: true,
+ input: map[string]any{
+ jsonKeyServiceEnv: map[string]any{
+ "NB_MANAGEMENT_URL": "https://mgmt.example.com:443",
+ "NB_LOG_LEVEL": "debug",
+ "NB_SETUP_KEY": "secret",
+ "SOME_OTHER": "val",
+ },
+ },
+ check: func(t *testing.T, params map[string]any) {
+ t.Helper()
+ env := params[jsonKeyServiceEnv].(map[string]any)
+ // Safe NB_ values should be anonymized (not the original, not masked)
+ mgmtVal := env["NB_MANAGEMENT_URL"].(string)
+ assert.NotEqual(t, "https://mgmt.example.com:443", mgmtVal, "should be anonymized")
+ assert.NotEqual(t, maskedValue, mgmtVal, "should not be masked")
+
+ logVal := env["NB_LOG_LEVEL"].(string)
+ assert.NotEqual(t, maskedValue, logVal, "safe NB_ var should not be masked")
+
+ // Sensitive and non-NB_ still masked
+ assert.Equal(t, maskedValue, env["NB_SETUP_KEY"])
+ assert.Equal(t, maskedValue, env["SOME_OTHER"])
+ },
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses())
+ g := &BundleGenerator{
+ anonymize: tt.anonymize,
+ anonymizer: anonymizer,
+ }
+ g.sanitizeServiceEnvVars(tt.input)
+ tt.check(t, tt.input)
+ })
+ }
+}
+
+func TestAddServiceParams(t *testing.T) {
+ t.Run("missing service.json returns nil", func(t *testing.T) {
+ g := &BundleGenerator{
+ anonymizer: anonymize.NewAnonymizer(anonymize.DefaultAddresses()),
+ }
+
+ origStateDir := configs.StateDir
+ configs.StateDir = t.TempDir()
+ t.Cleanup(func() { configs.StateDir = origStateDir })
+
+ err := g.addServiceParams()
+ assert.NoError(t, err)
+ })
+
+ t.Run("management_url anonymized when anonymize is true", func(t *testing.T) {
+ dir := t.TempDir()
+ origStateDir := configs.StateDir
+ configs.StateDir = dir
+ t.Cleanup(func() { configs.StateDir = origStateDir })
+
+ input := map[string]any{
+ jsonKeyManagementURL: "https://api.example.com:443",
+ jsonKeyServiceEnv: map[string]any{
+ "NB_LOG_LEVEL": "trace",
+ },
+ }
+ data, err := json.Marshal(input)
+ require.NoError(t, err)
+ require.NoError(t, os.WriteFile(filepath.Join(dir, serviceParamsFile), data, 0600))
+
+ var buf bytes.Buffer
+ zw := zip.NewWriter(&buf)
+
+ g := &BundleGenerator{
+ anonymize: true,
+ anonymizer: anonymize.NewAnonymizer(anonymize.DefaultAddresses()),
+ archive: zw,
+ }
+
+ require.NoError(t, g.addServiceParams())
+ require.NoError(t, zw.Close())
+
+ zr, err := zip.NewReader(bytes.NewReader(buf.Bytes()), int64(buf.Len()))
+ require.NoError(t, err)
+ require.Len(t, zr.File, 1)
+ assert.Equal(t, serviceParamsBundle, zr.File[0].Name)
+
+ rc, err := zr.File[0].Open()
+ require.NoError(t, err)
+ defer rc.Close()
+
+ var result map[string]any
+ require.NoError(t, json.NewDecoder(rc).Decode(&result))
+
+ mgmt := result[jsonKeyManagementURL].(string)
+ assert.NotEqual(t, "https://api.example.com:443", mgmt, "management_url should be anonymized")
+ assert.NotEmpty(t, mgmt)
+
+ env := result[jsonKeyServiceEnv].(map[string]any)
+ assert.NotEqual(t, maskedValue, env["NB_LOG_LEVEL"], "safe NB_ var should not be masked")
+ })
+
+ t.Run("management_url preserved when anonymize is false", func(t *testing.T) {
+ dir := t.TempDir()
+ origStateDir := configs.StateDir
+ configs.StateDir = dir
+ t.Cleanup(func() { configs.StateDir = origStateDir })
+
+ input := map[string]any{
+ jsonKeyManagementURL: "https://api.example.com:443",
+ }
+ data, err := json.Marshal(input)
+ require.NoError(t, err)
+ require.NoError(t, os.WriteFile(filepath.Join(dir, serviceParamsFile), data, 0600))
+
+ var buf bytes.Buffer
+ zw := zip.NewWriter(&buf)
+
+ g := &BundleGenerator{
+ anonymize: false,
+ anonymizer: anonymize.NewAnonymizer(anonymize.DefaultAddresses()),
+ archive: zw,
+ }
+
+ require.NoError(t, g.addServiceParams())
+ require.NoError(t, zw.Close())
+
+ zr, err := zip.NewReader(bytes.NewReader(buf.Bytes()), int64(buf.Len()))
+ require.NoError(t, err)
+
+ rc, err := zr.File[0].Open()
+ require.NoError(t, err)
+ defer rc.Close()
+
+ var result map[string]any
+ require.NoError(t, json.NewDecoder(rc).Decode(&result))
+
+ assert.Equal(t, "https://api.example.com:443", result[jsonKeyManagementURL], "management_url should be preserved")
+ })
+}
+
// Helper function to check if IP is in CGNAT range
func isInCGNATRange(ip net.IP) bool {
cgnat := net.IPNet{
diff --git a/client/internal/portforward/manager.go b/client/internal/portforward/manager.go
index 179dc3c45..b0680160c 100644
--- a/client/internal/portforward/manager.go
+++ b/client/internal/portforward/manager.go
@@ -6,6 +6,7 @@ import (
"context"
"fmt"
"net"
+ "regexp"
"sync"
"time"
@@ -17,20 +18,30 @@ import (
const (
defaultMappingTTL = 2 * time.Hour
- renewalInterval = defaultMappingTTL / 2
healthCheckInterval = 1 * time.Minute
discoveryTimeout = 10 * time.Second
mappingDescription = "NetBird"
)
+// upnpErrPermanentLeaseOnly matches UPnP error 725 in SOAP fault XML,
+// allowing for whitespace/newlines between tags from different router firmware.
+var upnpErrPermanentLeaseOnly = regexp.MustCompile(`\s*725\s*`)
+
+// Mapping represents an active NAT port mapping.
type Mapping struct {
Protocol string
InternalPort uint16
ExternalPort uint16
ExternalIP net.IP
NATType string
+ // TTL is the lease duration. Zero means a permanent lease that never expires.
+ TTL time.Duration
}
+// TODO: persist mapping state for crash recovery cleanup of permanent leases.
+// Currently not done because State.Cleanup requires NAT gateway re-discovery,
+// which blocks startup for ~10s when no gateway is present (affects all clients).
+
type Manager struct {
cancel context.CancelFunc
@@ -46,6 +57,7 @@ type Manager struct {
mu sync.Mutex
}
+// NewManager creates a new port forwarding manager.
func NewManager() *Manager {
return &Manager{
stopCtx: make(chan context.Context, 1),
@@ -80,8 +92,7 @@ func (m *Manager) Start(ctx context.Context, wgPort uint16) {
gateway, mapping, err := m.setup(ctx)
if err != nil {
- log.Errorf("failed to setup NAT port mapping: %v", err)
-
+ log.Infof("port forwarding setup: %v", err)
return
}
@@ -89,7 +100,7 @@ func (m *Manager) Start(ctx context.Context, wgPort uint16) {
m.mapping = mapping
m.mappingLock.Unlock()
- m.renewLoop(ctx, gateway)
+ m.renewLoop(ctx, gateway, mapping.TTL)
select {
case cleanupCtx := <-m.stopCtx:
@@ -148,16 +159,14 @@ func (m *Manager) setup(ctx context.Context) (nat.NAT, *Mapping, error) {
gateway, err := discoverGateway(discoverCtx)
if err != nil {
- log.Infof("NAT gateway discovery failed: %v (port forwarding disabled)", err)
- return nil, nil, err
+ return nil, nil, fmt.Errorf("discover gateway: %w", err)
}
log.Infof("discovered NAT gateway: %s", gateway.Type())
mapping, err := m.createMapping(ctx, gateway)
if err != nil {
- log.Warnf("failed to create port mapping: %v", err)
- return nil, nil, err
+ return nil, nil, fmt.Errorf("create port mapping: %w", err)
}
return gateway, mapping, nil
}
@@ -166,9 +175,18 @@ func (m *Manager) createMapping(ctx context.Context, gateway nat.NAT) (*Mapping,
ctx, cancel := context.WithTimeout(ctx, 30*time.Second)
defer cancel()
- externalPort, err := gateway.AddPortMapping(ctx, "udp", int(m.wgPort), mappingDescription, defaultMappingTTL)
+ ttl := defaultMappingTTL
+ externalPort, err := gateway.AddPortMapping(ctx, "udp", int(m.wgPort), mappingDescription, ttl)
if err != nil {
- return nil, err
+ if !isPermanentLeaseRequired(err) {
+ return nil, err
+ }
+ log.Infof("gateway only supports permanent leases, retrying with indefinite duration")
+ ttl = 0
+ externalPort, err = gateway.AddPortMapping(ctx, "udp", int(m.wgPort), mappingDescription, ttl)
+ if err != nil {
+ return nil, err
+ }
}
externalIP, err := gateway.GetExternalAddress()
@@ -182,6 +200,7 @@ func (m *Manager) createMapping(ctx context.Context, gateway nat.NAT) (*Mapping,
ExternalPort: uint16(externalPort),
ExternalIP: externalIP,
NATType: gateway.Type(),
+ TTL: ttl,
}
log.Infof("created port mapping: %d -> %d via %s (external IP: %s)",
@@ -189,8 +208,15 @@ func (m *Manager) createMapping(ctx context.Context, gateway nat.NAT) (*Mapping,
return mapping, nil
}
-func (m *Manager) renewLoop(ctx context.Context, gateway nat.NAT) {
- renewTicker := time.NewTicker(renewalInterval)
+func (m *Manager) renewLoop(ctx context.Context, gateway nat.NAT, ttl time.Duration) {
+ if ttl == 0 {
+ // Permanent mappings don't expire, just wait for cancellation
+ // but still run health checks for PCP gateways.
+ m.permanentLeaseLoop(ctx, gateway)
+ return
+ }
+
+ renewTicker := time.NewTicker(ttl / 2)
healthTicker := time.NewTicker(healthCheckInterval)
defer renewTicker.Stop()
defer healthTicker.Stop()
@@ -206,12 +232,26 @@ func (m *Manager) renewLoop(ctx context.Context, gateway nat.NAT) {
}
case <-healthTicker.C:
if m.checkHealthAndRecreate(ctx, gateway) {
- renewTicker.Reset(renewalInterval)
+ renewTicker.Reset(ttl / 2)
}
}
}
}
+func (m *Manager) permanentLeaseLoop(ctx context.Context, gateway nat.NAT) {
+ healthTicker := time.NewTicker(healthCheckInterval)
+ defer healthTicker.Stop()
+
+ for {
+ select {
+ case <-ctx.Done():
+ return
+ case <-healthTicker.C:
+ m.checkHealthAndRecreate(ctx, gateway)
+ }
+ }
+}
+
func (m *Manager) checkHealthAndRecreate(ctx context.Context, gateway nat.NAT) bool {
if isHealthCheckDisabled() {
return false
@@ -255,7 +295,7 @@ func (m *Manager) renewMapping(ctx context.Context, gateway nat.NAT) error {
ctx, cancel := context.WithTimeout(ctx, 30*time.Second)
defer cancel()
- externalPort, err := gateway.AddPortMapping(ctx, m.mapping.Protocol, int(m.mapping.InternalPort), mappingDescription, defaultMappingTTL)
+ externalPort, err := gateway.AddPortMapping(ctx, m.mapping.Protocol, int(m.mapping.InternalPort), mappingDescription, m.mapping.TTL)
if err != nil {
return fmt.Errorf("add port mapping: %w", err)
}
@@ -296,3 +336,7 @@ func (m *Manager) startTearDown(ctx context.Context) {
}
}
+// isPermanentLeaseRequired checks if a UPnP error indicates the gateway only supports permanent leases (error 725).
+func isPermanentLeaseRequired(err error) bool {
+ return err != nil && upnpErrPermanentLeaseOnly.MatchString(err.Error())
+}
diff --git a/client/internal/portforward/manager_js.go b/client/internal/portforward/manager_js.go
index d5db147f2..36c55063b 100644
--- a/client/internal/portforward/manager_js.go
+++ b/client/internal/portforward/manager_js.go
@@ -3,15 +3,18 @@ package portforward
import (
"context"
"net"
+ "time"
)
-// Mapping represents port mapping information.
+// Mapping represents an active NAT port mapping.
type Mapping struct {
Protocol string
InternalPort uint16
ExternalPort uint16
ExternalIP net.IP
NATType string
+ // TTL is the lease duration. Zero means a permanent lease that never expires.
+ TTL time.Duration
}
// Manager is a stub for js/wasm builds where NAT-PMP/UPnP is not supported.
diff --git a/client/internal/portforward/manager_test.go b/client/internal/portforward/manager_test.go
index 1029e87f5..1f66f9ccd 100644
--- a/client/internal/portforward/manager_test.go
+++ b/client/internal/portforward/manager_test.go
@@ -4,23 +4,25 @@ package portforward
import (
"context"
+ "fmt"
"net"
"testing"
"time"
- "github.com/libp2p/go-nat"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
type mockNAT struct {
- natType string
- deviceAddr net.IP
- externalAddr net.IP
- internalAddr net.IP
- mappings map[int]int
- addMappingErr error
- deleteMappingErr error
+ natType string
+ deviceAddr net.IP
+ externalAddr net.IP
+ internalAddr net.IP
+ mappings map[int]int
+ addMappingErr error
+ deleteMappingErr error
+ onlyPermanentLeases bool
+ lastTimeout time.Duration
}
func newMockNAT() *mockNAT {
@@ -53,8 +55,12 @@ func (m *mockNAT) AddPortMapping(ctx context.Context, protocol string, internalP
if m.addMappingErr != nil {
return 0, m.addMappingErr
}
+ if m.onlyPermanentLeases && timeout != 0 {
+ return 0, fmt.Errorf("SOAP fault. Code: | Explanation: | Detail: 725OnlyPermanentLeasesSupported")
+ }
externalPort := internalPort
m.mappings[internalPort] = externalPort
+ m.lastTimeout = timeout
return externalPort, nil
}
@@ -80,6 +86,7 @@ func TestManager_CreateMapping(t *testing.T) {
assert.Equal(t, uint16(51820), mapping.ExternalPort)
assert.Equal(t, "Mock-NAT", mapping.NATType)
assert.Equal(t, net.ParseIP("203.0.113.50").To4(), mapping.ExternalIP.To4())
+ assert.Equal(t, defaultMappingTTL, mapping.TTL)
}
func TestManager_GetMapping_ReturnsNilWhenNotReady(t *testing.T) {
@@ -131,29 +138,64 @@ func TestManager_Cleanup_NilMapping(t *testing.T) {
m.cleanup(context.Background(), gateway)
}
-func TestState_Cleanup(t *testing.T) {
- origDiscover := discoverGateway
- defer func() { discoverGateway = origDiscover }()
- mockGateway := newMockNAT()
- mockGateway.mappings[51820] = 51820
- discoverGateway = func(ctx context.Context) (nat.NAT, error) {
- return mockGateway, nil
- }
+func TestManager_CreateMapping_PermanentLeaseFallback(t *testing.T) {
+ m := NewManager()
+ m.wgPort = 51820
- state := &State{
- Protocol: "udp",
- InternalPort: 51820,
- }
+ gateway := newMockNAT()
+ gateway.onlyPermanentLeases = true
- err := state.Cleanup()
- assert.NoError(t, err)
+ mapping, err := m.createMapping(context.Background(), gateway)
+ require.NoError(t, err)
+ require.NotNil(t, mapping)
- _, exists := mockGateway.mappings[51820]
- assert.False(t, exists, "mapping should be deleted after cleanup")
+ assert.Equal(t, uint16(51820), mapping.InternalPort)
+ assert.Equal(t, time.Duration(0), mapping.TTL, "should return zero TTL for permanent lease")
+ assert.Equal(t, time.Duration(0), gateway.lastTimeout, "should have retried with zero duration")
}
-func TestState_Name(t *testing.T) {
- state := &State{}
- assert.Equal(t, "port_forward_state", state.Name())
+func TestIsPermanentLeaseRequired(t *testing.T) {
+ tests := []struct {
+ name string
+ err error
+ expected bool
+ }{
+ {
+ name: "nil error",
+ err: nil,
+ expected: false,
+ },
+ {
+ name: "UPnP error 725",
+ err: fmt.Errorf("SOAP fault. Code: | Detail: 725OnlyPermanentLeasesSupported"),
+ expected: true,
+ },
+ {
+ name: "wrapped error with 725",
+ err: fmt.Errorf("add port mapping: %w", fmt.Errorf("Detail: 725")),
+ expected: true,
+ },
+ {
+ name: "error 725 with newlines in XML",
+ err: fmt.Errorf("\n 725\n"),
+ expected: true,
+ },
+ {
+ name: "bare 725 without XML tag",
+ err: fmt.Errorf("error code 725"),
+ expected: false,
+ },
+ {
+ name: "unrelated error",
+ err: fmt.Errorf("connection refused"),
+ expected: false,
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ assert.Equal(t, tt.expected, isPermanentLeaseRequired(tt.err))
+ })
+ }
}
diff --git a/client/internal/routemanager/manager.go b/client/internal/routemanager/manager.go
index e7ca44239..3923e153b 100644
--- a/client/internal/routemanager/manager.go
+++ b/client/internal/routemanager/manager.go
@@ -168,6 +168,7 @@ func (m *DefaultManager) setupAndroidRoutes(config ManagerConfig) {
NetworkType: route.IPv4Network,
}
cr = append(cr, fakeIPRoute)
+ m.notifier.SetFakeIPRoute(fakeIPRoute)
}
m.notifier.SetInitialClientRoutes(cr, routesForComparison)
diff --git a/client/internal/routemanager/notifier/notifier_android.go b/client/internal/routemanager/notifier/notifier_android.go
index 3d2784ae1..55e0b7421 100644
--- a/client/internal/routemanager/notifier/notifier_android.go
+++ b/client/internal/routemanager/notifier/notifier_android.go
@@ -16,6 +16,7 @@ import (
type Notifier struct {
initialRoutes []*route.Route
currentRoutes []*route.Route
+ fakeIPRoute *route.Route
listener listener.NetworkChangeListener
listenerMux sync.Mutex
@@ -31,13 +32,17 @@ func (n *Notifier) SetListener(listener listener.NetworkChangeListener) {
n.listener = listener
}
-// SetInitialClientRoutes stores the full initial route set (including fake IP blocks)
-// and a separate comparison set (without fake IP blocks) for diff detection.
+// SetInitialClientRoutes stores the initial route sets for TUN configuration.
func (n *Notifier) SetInitialClientRoutes(initialRoutes []*route.Route, routesForComparison []*route.Route) {
n.initialRoutes = filterStatic(initialRoutes)
n.currentRoutes = filterStatic(routesForComparison)
}
+// SetFakeIPRoute stores the fake IP route to be included in every TUN rebuild.
+func (n *Notifier) SetFakeIPRoute(r *route.Route) {
+ n.fakeIPRoute = r
+}
+
func (n *Notifier) OnNewRoutes(idMap route.HAMap) {
var newRoutes []*route.Route
for _, routes := range idMap {
@@ -69,7 +74,9 @@ func (n *Notifier) notify() {
}
allRoutes := slices.Clone(n.currentRoutes)
- allRoutes = append(allRoutes, n.extraInitialRoutes()...)
+ if n.fakeIPRoute != nil {
+ allRoutes = append(allRoutes, n.fakeIPRoute)
+ }
routeStrings := n.routesToStrings(allRoutes)
sort.Strings(routeStrings)
@@ -78,23 +85,6 @@ func (n *Notifier) notify() {
}(n.listener)
}
-// extraInitialRoutes returns initialRoutes whose network prefix is absent
-// from currentRoutes (e.g. the fake IP block added at setup time).
-func (n *Notifier) extraInitialRoutes() []*route.Route {
- currentNets := make(map[netip.Prefix]struct{}, len(n.currentRoutes))
- for _, r := range n.currentRoutes {
- currentNets[r.Network] = struct{}{}
- }
-
- var extra []*route.Route
- for _, r := range n.initialRoutes {
- if _, ok := currentNets[r.Network]; !ok {
- extra = append(extra, r)
- }
- }
- return extra
-}
-
func filterStatic(routes []*route.Route) []*route.Route {
out := make([]*route.Route, 0, len(routes))
for _, r := range routes {
diff --git a/client/internal/routemanager/notifier/notifier_ios.go b/client/internal/routemanager/notifier/notifier_ios.go
index 343d2799e..68c85067a 100644
--- a/client/internal/routemanager/notifier/notifier_ios.go
+++ b/client/internal/routemanager/notifier/notifier_ios.go
@@ -34,6 +34,10 @@ func (n *Notifier) SetInitialClientRoutes([]*route.Route, []*route.Route) {
// iOS doesn't care about initial routes
}
+func (n *Notifier) SetFakeIPRoute(*route.Route) {
+ // Not used on iOS
+}
+
func (n *Notifier) OnNewRoutes(route.HAMap) {
// Not used on iOS
}
diff --git a/client/internal/routemanager/notifier/notifier_other.go b/client/internal/routemanager/notifier/notifier_other.go
index 0521e3dc2..97c815cf0 100644
--- a/client/internal/routemanager/notifier/notifier_other.go
+++ b/client/internal/routemanager/notifier/notifier_other.go
@@ -23,6 +23,10 @@ func (n *Notifier) SetInitialClientRoutes([]*route.Route, []*route.Route) {
// Not used on non-mobile platforms
}
+func (n *Notifier) SetFakeIPRoute(*route.Route) {
+ // Not used on non-mobile platforms
+}
+
func (n *Notifier) OnNewRoutes(idMap route.HAMap) {
// Not used on non-mobile platforms
}
diff --git a/client/server/state_generic.go b/client/server/state_generic.go
index 3f794b611..86475ca42 100644
--- a/client/server/state_generic.go
+++ b/client/server/state_generic.go
@@ -10,10 +10,6 @@ import (
)
// registerStates registers all states that need crash recovery cleanup.
-// Note: portforward.State is intentionally NOT registered here to avoid blocking startup
-// for up to 10 seconds during NAT gateway discovery when no gateway is present.
-// The gateway reference cannot be persisted across restarts, so cleanup requires re-discovery.
-// Port forward cleanup is handled by the Manager during normal operation instead.
func registerStates(mgr *statemanager.Manager) {
mgr.RegisterState(&dns.ShutdownState{})
mgr.RegisterState(&systemops.ShutdownState{})
diff --git a/client/server/state_linux.go b/client/server/state_linux.go
index 655edfc53..b193d4dfa 100644
--- a/client/server/state_linux.go
+++ b/client/server/state_linux.go
@@ -12,10 +12,6 @@ import (
)
// registerStates registers all states that need crash recovery cleanup.
-// Note: portforward.State is intentionally NOT registered here to avoid blocking startup
-// for up to 10 seconds during NAT gateway discovery when no gateway is present.
-// The gateway reference cannot be persisted across restarts, so cleanup requires re-discovery.
-// Port forward cleanup is handled by the Manager during normal operation instead.
func registerStates(mgr *statemanager.Manager) {
mgr.RegisterState(&dns.ShutdownState{})
mgr.RegisterState(&systemops.ShutdownState{})
diff --git a/client/system/info_freebsd.go b/client/system/info_freebsd.go
index 8e1353151..755172842 100644
--- a/client/system/info_freebsd.go
+++ b/client/system/info_freebsd.go
@@ -43,18 +43,24 @@ func GetInfo(ctx context.Context) *Info {
systemHostname, _ := os.Hostname()
+ addrs, err := networkAddresses()
+ if err != nil {
+ log.Warnf("failed to discover network addresses: %s", err)
+ }
+
return &Info{
- GoOS: runtime.GOOS,
- Kernel: osInfo[0],
- Platform: runtime.GOARCH,
- OS: osName,
- OSVersion: osVersion,
- Hostname: extractDeviceName(ctx, systemHostname),
- CPUs: runtime.NumCPU(),
- NetbirdVersion: version.NetbirdVersion(),
- UIVersion: extractUserAgent(ctx),
- KernelVersion: osInfo[1],
- Environment: env,
+ GoOS: runtime.GOOS,
+ Kernel: osInfo[0],
+ Platform: runtime.GOARCH,
+ OS: osName,
+ OSVersion: osVersion,
+ Hostname: extractDeviceName(ctx, systemHostname),
+ CPUs: runtime.NumCPU(),
+ NetbirdVersion: version.NetbirdVersion(),
+ UIVersion: extractUserAgent(ctx),
+ KernelVersion: osInfo[1],
+ NetworkAddresses: addrs,
+ Environment: env,
}
}
diff --git a/client/ui/debug.go b/client/ui/debug.go
index 29f73a66a..4ebe4d675 100644
--- a/client/ui/debug.go
+++ b/client/ui/debug.go
@@ -24,9 +24,10 @@ import (
// Initial state for the debug collection
type debugInitialState struct {
- wasDown bool
- logLevel proto.LogLevel
- isLevelTrace bool
+ wasDown bool
+ needsRestoreUp bool
+ logLevel proto.LogLevel
+ isLevelTrace bool
}
// Debug collection parameters
@@ -371,46 +372,51 @@ func (s *serviceClient) configureServiceForDebug(
conn proto.DaemonServiceClient,
state *debugInitialState,
enablePersistence bool,
-) error {
+) {
if state.wasDown {
if _, err := conn.Up(s.ctx, &proto.UpRequest{}); err != nil {
- return fmt.Errorf("bring service up: %v", err)
+ log.Warnf("failed to bring service up: %v", err)
+ } else {
+ log.Info("Service brought up for debug")
+ time.Sleep(time.Second * 10)
}
- log.Info("Service brought up for debug")
- time.Sleep(time.Second * 10)
}
if !state.isLevelTrace {
if _, err := conn.SetLogLevel(s.ctx, &proto.SetLogLevelRequest{Level: proto.LogLevel_TRACE}); err != nil {
- return fmt.Errorf("set log level to TRACE: %v", err)
+ log.Warnf("failed to set log level to TRACE: %v", err)
+ } else {
+ log.Info("Log level set to TRACE for debug")
}
- log.Info("Log level set to TRACE for debug")
}
if _, err := conn.Down(s.ctx, &proto.DownRequest{}); err != nil {
- return fmt.Errorf("bring service down: %v", err)
+ log.Warnf("failed to bring service down: %v", err)
+ } else {
+ state.needsRestoreUp = !state.wasDown
+ time.Sleep(time.Second)
}
- time.Sleep(time.Second)
if enablePersistence {
if _, err := conn.SetSyncResponsePersistence(s.ctx, &proto.SetSyncResponsePersistenceRequest{
Enabled: true,
}); err != nil {
- return fmt.Errorf("enable sync response persistence: %v", err)
+ log.Warnf("failed to enable sync response persistence: %v", err)
+ } else {
+ log.Info("Sync response persistence enabled for debug")
}
- log.Info("Sync response persistence enabled for debug")
}
if _, err := conn.Up(s.ctx, &proto.UpRequest{}); err != nil {
- return fmt.Errorf("bring service back up: %v", err)
+ log.Warnf("failed to bring service back up: %v", err)
+ } else {
+ state.needsRestoreUp = false
+ time.Sleep(time.Second * 3)
}
- time.Sleep(time.Second * 3)
if _, err := conn.StartCPUProfile(s.ctx, &proto.StartCPUProfileRequest{}); err != nil {
log.Warnf("failed to start CPU profiling: %v", err)
}
-
- return nil
}
func (s *serviceClient) collectDebugData(
@@ -424,9 +430,7 @@ func (s *serviceClient) collectDebugData(
var wg sync.WaitGroup
startProgressTracker(ctx, &wg, params.duration, progress)
- if err := s.configureServiceForDebug(conn, state, params.enablePersistence); err != nil {
- return err
- }
+ s.configureServiceForDebug(conn, state, params.enablePersistence)
wg.Wait()
progress.progressBar.Hide()
@@ -482,9 +486,17 @@ func (s *serviceClient) createDebugBundleFromCollection(
// Restore service to original state
func (s *serviceClient) restoreServiceState(conn proto.DaemonServiceClient, state *debugInitialState) {
+ if state.needsRestoreUp {
+ if _, err := conn.Up(s.ctx, &proto.UpRequest{}); err != nil {
+ log.Warnf("failed to restore up state: %v", err)
+ } else {
+ log.Info("Service state restored to up")
+ }
+ }
+
if state.wasDown {
if _, err := conn.Down(s.ctx, &proto.DownRequest{}); err != nil {
- log.Errorf("Failed to restore down state: %v", err)
+ log.Warnf("failed to restore down state: %v", err)
} else {
log.Info("Service state restored to down")
}
@@ -492,7 +504,7 @@ func (s *serviceClient) restoreServiceState(conn proto.DaemonServiceClient, stat
if !state.isLevelTrace {
if _, err := conn.SetLogLevel(s.ctx, &proto.SetLogLevelRequest{Level: state.logLevel}); err != nil {
- log.Errorf("Failed to restore log level: %v", err)
+ log.Warnf("failed to restore log level: %v", err)
} else {
log.Info("Log level restored to original setting")
}
diff --git a/combined/cmd/config.go b/combined/cmd/config.go
index 85664d0d2..ce4df8394 100644
--- a/combined/cmd/config.go
+++ b/combined/cmd/config.go
@@ -179,9 +179,11 @@ type StoreConfig struct {
// ReverseProxyConfig contains reverse proxy settings
type ReverseProxyConfig struct {
- TrustedHTTPProxies []string `yaml:"trustedHTTPProxies"`
- TrustedHTTPProxiesCount uint `yaml:"trustedHTTPProxiesCount"`
- TrustedPeers []string `yaml:"trustedPeers"`
+ TrustedHTTPProxies []string `yaml:"trustedHTTPProxies"`
+ TrustedHTTPProxiesCount uint `yaml:"trustedHTTPProxiesCount"`
+ TrustedPeers []string `yaml:"trustedPeers"`
+ AccessLogRetentionDays int `yaml:"accessLogRetentionDays"`
+ AccessLogCleanupIntervalHours int `yaml:"accessLogCleanupIntervalHours"`
}
// DefaultConfig returns a CombinedConfig with default values
@@ -645,7 +647,9 @@ func (c *CombinedConfig) ToManagementConfig() (*nbconfig.Config, error) {
// Build reverse proxy config
reverseProxy := nbconfig.ReverseProxy{
- TrustedHTTPProxiesCount: mgmt.ReverseProxy.TrustedHTTPProxiesCount,
+ TrustedHTTPProxiesCount: mgmt.ReverseProxy.TrustedHTTPProxiesCount,
+ AccessLogRetentionDays: mgmt.ReverseProxy.AccessLogRetentionDays,
+ AccessLogCleanupIntervalHours: mgmt.ReverseProxy.AccessLogCleanupIntervalHours,
}
for _, p := range mgmt.ReverseProxy.TrustedHTTPProxies {
if prefix, err := netip.ParsePrefix(p); err == nil {
diff --git a/flow/client/client.go b/flow/client/client.go
index 318fcfe1e..8ad637974 100644
--- a/flow/client/client.go
+++ b/flow/client/client.go
@@ -14,7 +14,6 @@ import (
log "github.com/sirupsen/logrus"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
- "google.golang.org/grpc/connectivity"
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/credentials/insecure"
"google.golang.org/grpc/keepalive"
@@ -26,11 +25,22 @@ import (
"github.com/netbirdio/netbird/util/wsproxy"
)
+var ErrClientClosed = errors.New("client is closed")
+
+// minHealthyDuration is the minimum time a stream must survive before a failure
+// resets the backoff timer. Streams that fail faster are considered unhealthy and
+// should not reset backoff, so that MaxElapsedTime can eventually stop retries.
+const minHealthyDuration = 5 * time.Second
+
type GRPCClient struct {
realClient proto.FlowServiceClient
clientConn *grpc.ClientConn
stream proto.FlowService_EventsClient
- streamMu sync.Mutex
+ target string
+ opts []grpc.DialOption
+ closed bool // prevent creating conn in the middle of the Close
+ receiving bool // prevent concurrent Receive calls
+ mu sync.Mutex // protects clientConn, realClient, stream, closed, and receiving
}
func NewClient(addr, payload, signature string, interval time.Duration) (*GRPCClient, error) {
@@ -65,7 +75,8 @@ func NewClient(addr, payload, signature string, interval time.Duration) (*GRPCCl
grpc.WithDefaultServiceConfig(`{"healthCheckConfig": {"serviceName": ""}}`),
)
- conn, err := grpc.NewClient(fmt.Sprintf("%s:%s", parsedURL.Hostname(), parsedURL.Port()), opts...)
+ target := parsedURL.Host
+ conn, err := grpc.NewClient(target, opts...)
if err != nil {
return nil, fmt.Errorf("creating new grpc client: %w", err)
}
@@ -73,30 +84,73 @@ func NewClient(addr, payload, signature string, interval time.Duration) (*GRPCCl
return &GRPCClient{
realClient: proto.NewFlowServiceClient(conn),
clientConn: conn,
+ target: target,
+ opts: opts,
}, nil
}
func (c *GRPCClient) Close() error {
- c.streamMu.Lock()
- defer c.streamMu.Unlock()
-
+ c.mu.Lock()
+ c.closed = true
c.stream = nil
- if err := c.clientConn.Close(); err != nil && !errors.Is(err, context.Canceled) {
+ conn := c.clientConn
+ c.clientConn = nil
+ c.mu.Unlock()
+
+ if conn == nil {
+ return nil
+ }
+
+ if err := conn.Close(); err != nil && !errors.Is(err, context.Canceled) {
return fmt.Errorf("close client connection: %w", err)
}
return nil
}
+func (c *GRPCClient) Send(event *proto.FlowEvent) error {
+ c.mu.Lock()
+ stream := c.stream
+ c.mu.Unlock()
+
+ if stream == nil {
+ return errors.New("stream not initialized")
+ }
+
+ if err := stream.Send(event); err != nil {
+ return fmt.Errorf("send flow event: %w", err)
+ }
+
+ return nil
+}
+
func (c *GRPCClient) Receive(ctx context.Context, interval time.Duration, msgHandler func(msg *proto.FlowEventAck) error) error {
+ c.mu.Lock()
+ if c.receiving {
+ c.mu.Unlock()
+ return errors.New("concurrent Receive calls are not supported")
+ }
+ c.receiving = true
+ c.mu.Unlock()
+ defer func() {
+ c.mu.Lock()
+ c.receiving = false
+ c.mu.Unlock()
+ }()
+
backOff := defaultBackoff(ctx, interval)
operation := func() error {
- if err := c.establishStreamAndReceive(ctx, msgHandler); err != nil {
- if s, ok := status.FromError(err); ok && s.Code() == codes.Canceled {
- return fmt.Errorf("receive: %w: %w", err, context.Canceled)
- }
+ stream, err := c.establishStream(ctx)
+ if err != nil {
+ log.Errorf("failed to establish flow stream, retrying: %v", err)
+ return c.handleRetryableError(err, time.Time{}, backOff)
+ }
+
+ streamStart := time.Now()
+
+ if err := c.receive(stream, msgHandler); err != nil {
log.Errorf("receive failed: %v", err)
- return fmt.Errorf("receive: %w", err)
+ return c.handleRetryableError(err, streamStart, backOff)
}
return nil
}
@@ -108,37 +162,106 @@ func (c *GRPCClient) Receive(ctx context.Context, interval time.Duration, msgHan
return nil
}
-func (c *GRPCClient) establishStreamAndReceive(ctx context.Context, msgHandler func(msg *proto.FlowEventAck) error) error {
- if c.clientConn.GetState() == connectivity.Shutdown {
- return errors.New("connection to flow receiver has been shut down")
+// handleRetryableError resets the backoff timer if the stream was healthy long
+// enough and recreates the underlying ClientConn so that gRPC's internal
+// subchannel backoff does not accumulate and compete with our own retry timer.
+// A zero streamStart means the stream was never established.
+func (c *GRPCClient) handleRetryableError(err error, streamStart time.Time, backOff backoff.BackOff) error {
+ if isContextDone(err) {
+ return backoff.Permanent(err)
}
- stream, err := c.realClient.Events(ctx, grpc.WaitForReady(true))
- if err != nil {
- return fmt.Errorf("create event stream: %w", err)
+ var permErr *backoff.PermanentError
+ if errors.As(err, &permErr) {
+ return err
}
- err = stream.Send(&proto.FlowEvent{IsInitiator: true})
+ // Reset the backoff so the next retry starts with a short delay instead of
+ // continuing the already-elapsed timer. Only do this if the stream was healthy
+ // long enough; short-lived connect/drop cycles must not defeat MaxElapsedTime.
+ if !streamStart.IsZero() && time.Since(streamStart) >= minHealthyDuration {
+ backOff.Reset()
+ }
+
+ if recreateErr := c.recreateConnection(); recreateErr != nil {
+ log.Errorf("recreate connection: %v", recreateErr)
+ return recreateErr
+ }
+
+ log.Infof("connection recreated, retrying stream")
+ return fmt.Errorf("retrying after error: %w", err)
+}
+
+func (c *GRPCClient) recreateConnection() error {
+ c.mu.Lock()
+ if c.closed {
+ c.mu.Unlock()
+ return backoff.Permanent(ErrClientClosed)
+ }
+
+ conn, err := grpc.NewClient(c.target, c.opts...)
if err != nil {
- log.Infof("failed to send initiator message to flow receiver but will attempt to continue. Error: %s", err)
+ c.mu.Unlock()
+ return fmt.Errorf("create new connection: %w", err)
+ }
+
+ old := c.clientConn
+ c.clientConn = conn
+ c.realClient = proto.NewFlowServiceClient(conn)
+ c.stream = nil
+ c.mu.Unlock()
+
+ _ = old.Close()
+
+ return nil
+}
+
+func (c *GRPCClient) establishStream(ctx context.Context) (proto.FlowService_EventsClient, error) {
+ c.mu.Lock()
+ if c.closed {
+ c.mu.Unlock()
+ return nil, backoff.Permanent(ErrClientClosed)
+ }
+ cl := c.realClient
+ c.mu.Unlock()
+
+ // open stream outside the lock — blocking operation
+ stream, err := cl.Events(ctx)
+ if err != nil {
+ return nil, fmt.Errorf("create event stream: %w", err)
+ }
+ streamReady := false
+ defer func() {
+ if !streamReady {
+ _ = stream.CloseSend()
+ }
+ }()
+
+ if err = stream.Send(&proto.FlowEvent{IsInitiator: true}); err != nil {
+ return nil, fmt.Errorf("send initiator: %w", err)
}
if err = checkHeader(stream); err != nil {
- return fmt.Errorf("check header: %w", err)
+ return nil, fmt.Errorf("check header: %w", err)
}
- c.streamMu.Lock()
+ c.mu.Lock()
+ if c.closed {
+ c.mu.Unlock()
+ return nil, backoff.Permanent(ErrClientClosed)
+ }
c.stream = stream
- c.streamMu.Unlock()
+ c.mu.Unlock()
+ streamReady = true
- return c.receive(stream, msgHandler)
+ return stream, nil
}
func (c *GRPCClient) receive(stream proto.FlowService_EventsClient, msgHandler func(msg *proto.FlowEventAck) error) error {
for {
msg, err := stream.Recv()
if err != nil {
- return fmt.Errorf("receive from stream: %w", err)
+ return err
}
if msg.IsInitiator {
@@ -169,7 +292,7 @@ func checkHeader(stream proto.FlowService_EventsClient) error {
func defaultBackoff(ctx context.Context, interval time.Duration) backoff.BackOff {
return backoff.WithContext(&backoff.ExponentialBackOff{
InitialInterval: 800 * time.Millisecond,
- RandomizationFactor: 1,
+ RandomizationFactor: 0.5,
Multiplier: 1.7,
MaxInterval: interval / 2,
MaxElapsedTime: 3 * 30 * 24 * time.Hour, // 3 months
@@ -178,18 +301,12 @@ func defaultBackoff(ctx context.Context, interval time.Duration) backoff.BackOff
}, ctx)
}
-func (c *GRPCClient) Send(event *proto.FlowEvent) error {
- c.streamMu.Lock()
- stream := c.stream
- c.streamMu.Unlock()
-
- if stream == nil {
- return errors.New("stream not initialized")
+func isContextDone(err error) bool {
+ if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
+ return true
}
-
- if err := stream.Send(event); err != nil {
- return fmt.Errorf("send flow event: %w", err)
+ if s, ok := status.FromError(err); ok {
+ return s.Code() == codes.Canceled || s.Code() == codes.DeadlineExceeded
}
-
- return nil
+ return false
}
diff --git a/flow/client/client_test.go b/flow/client/client_test.go
index efe01c003..55157acbc 100644
--- a/flow/client/client_test.go
+++ b/flow/client/client_test.go
@@ -2,8 +2,11 @@ package client_test
import (
"context"
+ "encoding/binary"
"errors"
"net"
+ "sync"
+ "sync/atomic"
"testing"
"time"
@@ -11,6 +14,8 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/grpc"
+ "google.golang.org/grpc/codes"
+ "google.golang.org/grpc/status"
flow "github.com/netbirdio/netbird/flow/client"
"github.com/netbirdio/netbird/flow/proto"
@@ -18,21 +23,89 @@ import (
type testServer struct {
proto.UnimplementedFlowServiceServer
- events chan *proto.FlowEvent
- acks chan *proto.FlowEventAck
- grpcSrv *grpc.Server
- addr string
+ events chan *proto.FlowEvent
+ acks chan *proto.FlowEventAck
+ grpcSrv *grpc.Server
+ addr string
+ listener *connTrackListener
+ closeStream chan struct{} // signal server to close the stream
+ handlerDone chan struct{} // signaled each time Events() exits
+ handlerStarted chan struct{} // signaled each time Events() begins
+}
+
+// connTrackListener wraps a net.Listener to track accepted connections
+// so tests can forcefully close them to simulate PROTOCOL_ERROR/RST_STREAM.
+type connTrackListener struct {
+ net.Listener
+ mu sync.Mutex
+ conns []net.Conn
+}
+
+func (l *connTrackListener) Accept() (net.Conn, error) {
+ c, err := l.Listener.Accept()
+ if err != nil {
+ return nil, err
+ }
+ l.mu.Lock()
+ l.conns = append(l.conns, c)
+ l.mu.Unlock()
+ return c, nil
+}
+
+// sendRSTStream writes a raw HTTP/2 RST_STREAM frame with PROTOCOL_ERROR
+// (error code 0x1) on every tracked connection. This produces the exact error:
+//
+// rpc error: code = Internal desc = stream terminated by RST_STREAM with error code: PROTOCOL_ERROR
+//
+// HTTP/2 RST_STREAM frame format (9-byte header + 4-byte payload):
+//
+// Length (3 bytes): 0x000004
+// Type (1 byte): 0x03 (RST_STREAM)
+// Flags (1 byte): 0x00
+// Stream ID (4 bytes): target stream (must have bit 31 clear)
+// Error Code (4 bytes): 0x00000001 (PROTOCOL_ERROR)
+func (l *connTrackListener) connCount() int {
+ l.mu.Lock()
+ defer l.mu.Unlock()
+ return len(l.conns)
+}
+
+func (l *connTrackListener) sendRSTStream(streamID uint32) {
+ l.mu.Lock()
+ defer l.mu.Unlock()
+
+ frame := make([]byte, 13) // 9-byte header + 4-byte payload
+ // Length = 4 (3 bytes, big-endian)
+ frame[0], frame[1], frame[2] = 0, 0, 4
+ // Type = RST_STREAM (0x03)
+ frame[3] = 0x03
+ // Flags = 0
+ frame[4] = 0x00
+ // Stream ID (4 bytes, big-endian, bit 31 reserved = 0)
+ binary.BigEndian.PutUint32(frame[5:9], streamID)
+ // Error Code = PROTOCOL_ERROR (0x1)
+ binary.BigEndian.PutUint32(frame[9:13], 0x1)
+
+ for _, c := range l.conns {
+ _, _ = c.Write(frame)
+ }
}
func newTestServer(t *testing.T) *testServer {
- listener, err := net.Listen("tcp", "127.0.0.1:0")
+ rawListener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
+ listener := &connTrackListener{Listener: rawListener}
+
s := &testServer{
- events: make(chan *proto.FlowEvent, 100),
- acks: make(chan *proto.FlowEventAck, 100),
- grpcSrv: grpc.NewServer(),
- addr: listener.Addr().String(),
+ events: make(chan *proto.FlowEvent, 100),
+ acks: make(chan *proto.FlowEventAck, 100),
+ grpcSrv: grpc.NewServer(),
+ addr: rawListener.Addr().String(),
+ listener: listener,
+ closeStream: make(chan struct{}, 1),
+ handlerDone: make(chan struct{}, 10),
+ handlerStarted: make(chan struct{}, 10),
}
proto.RegisterFlowServiceServer(s.grpcSrv, s)
@@ -51,11 +124,23 @@ func newTestServer(t *testing.T) *testServer {
}
func (s *testServer) Events(stream proto.FlowService_EventsServer) error {
+ defer func() {
+ select {
+ case s.handlerDone <- struct{}{}:
+ default:
+ }
+ }()
+
err := stream.Send(&proto.FlowEventAck{IsInitiator: true})
if err != nil {
return err
}
+ select {
+ case s.handlerStarted <- struct{}{}:
+ default:
+ }
+
ctx, cancel := context.WithCancel(stream.Context())
defer cancel()
@@ -91,6 +176,8 @@ func (s *testServer) Events(stream proto.FlowService_EventsServer) error {
if err := stream.Send(ack); err != nil {
return err
}
+ case <-s.closeStream:
+ return status.Errorf(codes.Internal, "server closing stream")
case <-ctx.Done():
return ctx.Err()
}
@@ -110,16 +197,13 @@ func TestReceive(t *testing.T) {
assert.NoError(t, err, "failed to close flow")
})
- receivedAcks := make(map[string]bool)
+ var ackCount atomic.Int32
receiveDone := make(chan struct{})
go func() {
err := client.Receive(ctx, 1*time.Second, func(msg *proto.FlowEventAck) error {
if !msg.IsInitiator && len(msg.EventId) > 0 {
- id := string(msg.EventId)
- receivedAcks[id] = true
-
- if len(receivedAcks) >= 3 {
+ if ackCount.Add(1) >= 3 {
close(receiveDone)
}
}
@@ -130,7 +214,11 @@ func TestReceive(t *testing.T) {
}
}()
- time.Sleep(500 * time.Millisecond)
+ select {
+ case <-server.handlerStarted:
+ case <-time.After(3 * time.Second):
+ t.Fatal("timeout waiting for stream to be established")
+ }
for i := 0; i < 3; i++ {
eventID := uuid.New().String()
@@ -153,7 +241,7 @@ func TestReceive(t *testing.T) {
t.Fatal("timeout waiting for acks to be processed")
}
- assert.Equal(t, 3, len(receivedAcks))
+ assert.Equal(t, int32(3), ackCount.Load())
}
func TestReceive_ContextCancellation(t *testing.T) {
@@ -254,3 +342,195 @@ func TestSend(t *testing.T) {
t.Fatal("timeout waiting for ack to be received by flow")
}
}
+
+func TestNewClient_PermanentClose(t *testing.T) {
+ server := newTestServer(t)
+
+ client, err := flow.NewClient("http://"+server.addr, "test-payload", "test-signature", 1*time.Second)
+ require.NoError(t, err)
+
+ err = client.Close()
+ require.NoError(t, err)
+
+ ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
+ t.Cleanup(cancel)
+
+ done := make(chan error, 1)
+ go func() {
+ done <- client.Receive(ctx, 1*time.Second, func(msg *proto.FlowEventAck) error {
+ return nil
+ })
+ }()
+
+ select {
+ case err := <-done:
+ require.ErrorIs(t, err, flow.ErrClientClosed)
+ case <-time.After(2 * time.Second):
+ t.Fatal("Receive did not return after Close — stuck in retry loop")
+ }
+}
+
+func TestNewClient_CloseVerify(t *testing.T) {
+ server := newTestServer(t)
+
+ client, err := flow.NewClient("http://"+server.addr, "test-payload", "test-signature", 1*time.Second)
+ require.NoError(t, err)
+
+ ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
+ t.Cleanup(cancel)
+
+ done := make(chan error, 1)
+ go func() {
+ done <- client.Receive(ctx, 1*time.Second, func(msg *proto.FlowEventAck) error {
+ return nil
+ })
+ }()
+
+ closeDone := make(chan struct{}, 1)
+ go func() {
+ _ = client.Close()
+ closeDone <- struct{}{}
+ }()
+
+ select {
+ case err := <-done:
+ require.Error(t, err)
+ case <-time.After(2 * time.Second):
+ t.Fatal("Receive did not return after Close — stuck in retry loop")
+ }
+
+ select {
+ case <-closeDone:
+ return
+ case <-time.After(2 * time.Second):
+ t.Fatal("Close did not return — blocked in retry loop")
+ }
+
+}
+
+func TestClose_WhileReceiving(t *testing.T) {
+ server := newTestServer(t)
+ client, err := flow.NewClient("http://"+server.addr, "test-payload", "test-signature", 1*time.Second)
+ require.NoError(t, err)
+
+ ctx := context.Background() // no timeout — intentional
+ receiveDone := make(chan struct{})
+ go func() {
+ _ = client.Receive(ctx, 1*time.Second, func(msg *proto.FlowEventAck) error {
+ return nil
+ })
+ close(receiveDone)
+ }()
+
+ // Wait for the server-side handler to confirm the stream is established.
+ select {
+ case <-server.handlerStarted:
+ case <-time.After(3 * time.Second):
+ t.Fatal("timeout waiting for stream to be established")
+ }
+
+ closeDone := make(chan struct{})
+ go func() {
+ _ = client.Close()
+ close(closeDone)
+ }()
+
+ select {
+ case <-closeDone:
+ // Close returned — good
+ case <-time.After(2 * time.Second):
+ t.Fatal("Close blocked forever — Receive stuck in retry loop")
+ }
+
+ select {
+ case <-receiveDone:
+ case <-time.After(2 * time.Second):
+ t.Fatal("Receive did not exit after Close")
+ }
+}
+
+func TestReceive_ProtocolErrorStreamReconnect(t *testing.T) {
+ server := newTestServer(t)
+
+ ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
+ t.Cleanup(cancel)
+
+ client, err := flow.NewClient("http://"+server.addr, "test-payload", "test-signature", 1*time.Second)
+ require.NoError(t, err)
+ t.Cleanup(func() {
+ err := client.Close()
+ assert.NoError(t, err, "failed to close flow")
+ })
+
+ // Track acks received before and after server-side stream close
+ var ackCount atomic.Int32
+ receivedFirst := make(chan struct{})
+ receivedAfterReconnect := make(chan struct{})
+
+ go func() {
+ err := client.Receive(ctx, 1*time.Second, func(msg *proto.FlowEventAck) error {
+ if msg.IsInitiator || len(msg.EventId) == 0 {
+ return nil
+ }
+ n := ackCount.Add(1)
+ if n == 1 {
+ close(receivedFirst)
+ }
+ if n == 2 {
+ close(receivedAfterReconnect)
+ }
+ return nil
+ })
+ if err != nil && !errors.Is(err, context.Canceled) {
+ t.Logf("receive error: %v", err)
+ }
+ }()
+
+ // Wait for stream to be established, then send first ack
+ select {
+ case <-server.handlerStarted:
+ case <-time.After(3 * time.Second):
+ t.Fatal("timeout waiting for stream to be established")
+ }
+ server.acks <- &proto.FlowEventAck{EventId: []byte("before-close")}
+
+ select {
+ case <-receivedFirst:
+ case <-time.After(3 * time.Second):
+ t.Fatal("timeout waiting for first ack")
+ }
+
+ // Snapshot connection count before injecting the fault.
+ connsBefore := server.listener.connCount()
+
+ // Send a raw HTTP/2 RST_STREAM frame with PROTOCOL_ERROR on the TCP connection.
+ // gRPC multiplexes streams on stream IDs 1, 3, 5, ... (odd, client-initiated).
+ // Stream ID 1 is the client's first stream (our Events bidi stream).
+ // This produces the exact error the client sees in production:
+ // "stream terminated by RST_STREAM with error code: PROTOCOL_ERROR"
+ server.listener.sendRSTStream(1)
+
+ // Wait for the old Events() handler to fully exit so it can no longer
+ // drain s.acks and drop our injected ack on a broken stream.
+ select {
+ case <-server.handlerDone:
+ case <-time.After(5 * time.Second):
+ t.Fatal("old Events() handler did not exit after RST_STREAM")
+ }
+
+ require.Eventually(t, func() bool {
+ return server.listener.connCount() > connsBefore
+ }, 5*time.Second, 50*time.Millisecond, "client did not open a new TCP connection after RST_STREAM")
+
+ server.acks <- &proto.FlowEventAck{EventId: []byte("after-close")}
+
+ select {
+ case <-receivedAfterReconnect:
+ // Client successfully reconnected and received ack after server-side stream close
+ case <-time.After(5 * time.Second):
+ t.Fatal("timeout waiting for ack after server-side stream close — client did not reconnect")
+ }
+
+ assert.GreaterOrEqual(t, int(ackCount.Load()), 2, "should have received acks before and after stream close")
+ assert.GreaterOrEqual(t, server.listener.connCount(), 2, "client should have created at least 2 TCP connections (original + reconnect)")
+}
diff --git a/infrastructure_files/observability/grafana/dashboards/management.json b/infrastructure_files/observability/grafana/dashboards/management.json
index 95983603f..f116a8bde 100644
--- a/infrastructure_files/observability/grafana/dashboards/management.json
+++ b/infrastructure_files/observability/grafana/dashboards/management.json
@@ -302,7 +302,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "rate(management_account_peer_meta_update_counter_ratio_total{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])",
+ "expr": "rate(management_account_peer_meta_update_counter_total{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])",
"instant": false,
"legendFormat": "{{cluster}}/{{environment}}/{{job}}",
"range": true,
@@ -410,7 +410,7 @@
},
"disableTextWrap": false,
"editorMode": "code",
- "expr": "histogram_quantile(0.5,sum(increase(management_account_get_peer_network_map_duration_ms_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
+ "expr": "histogram_quantile(0.5,sum(increase(management_account_get_peer_network_map_duration_ms_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
"format": "heatmap",
"fullMetaSearch": false,
"includeNullMetadata": true,
@@ -426,7 +426,7 @@
},
"disableTextWrap": false,
"editorMode": "code",
- "expr": "histogram_quantile(0.9,sum(increase(management_account_get_peer_network_map_duration_ms_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
+ "expr": "histogram_quantile(0.9,sum(increase(management_account_get_peer_network_map_duration_ms_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
"format": "heatmap",
"fullMetaSearch": false,
"hide": false,
@@ -443,7 +443,7 @@
},
"disableTextWrap": false,
"editorMode": "code",
- "expr": "histogram_quantile(0.99,sum(increase(management_account_get_peer_network_map_duration_ms_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
+ "expr": "histogram_quantile(0.99,sum(increase(management_account_get_peer_network_map_duration_ms_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
"format": "heatmap",
"fullMetaSearch": false,
"hide": false,
@@ -545,7 +545,7 @@
},
"disableTextWrap": false,
"editorMode": "code",
- "expr": "histogram_quantile(0.5,sum(increase(management_account_update_account_peers_duration_ms_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
+ "expr": "histogram_quantile(0.5,sum(increase(management_account_update_account_peers_duration_ms_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
"format": "heatmap",
"fullMetaSearch": false,
"includeNullMetadata": true,
@@ -561,7 +561,7 @@
},
"disableTextWrap": false,
"editorMode": "code",
- "expr": "histogram_quantile(0.9,sum(increase(management_account_update_account_peers_duration_ms_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
+ "expr": "histogram_quantile(0.9,sum(increase(management_account_update_account_peers_duration_ms_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
"format": "heatmap",
"fullMetaSearch": false,
"hide": false,
@@ -578,7 +578,7 @@
},
"disableTextWrap": false,
"editorMode": "code",
- "expr": "histogram_quantile(0.99,sum(increase(management_account_update_account_peers_duration_ms_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
+ "expr": "histogram_quantile(0.99,sum(increase(management_account_update_account_peers_duration_ms_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
"format": "heatmap",
"fullMetaSearch": false,
"hide": false,
@@ -694,7 +694,7 @@
},
"disableTextWrap": false,
"editorMode": "code",
- "expr": "histogram_quantile(0.5,sum(increase(management_grpc_updatechannel_queue_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
+ "expr": "histogram_quantile(0.5,sum(increase(management_grpc_updatechannel_queue_length_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
"format": "heatmap",
"fullMetaSearch": false,
"includeNullMetadata": true,
@@ -710,7 +710,7 @@
},
"disableTextWrap": false,
"editorMode": "code",
- "expr": "histogram_quantile(0.9,sum(increase(management_grpc_updatechannel_queue_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
+ "expr": "histogram_quantile(0.9,sum(increase(management_grpc_updatechannel_queue_length_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
"format": "heatmap",
"fullMetaSearch": false,
"hide": false,
@@ -727,7 +727,7 @@
},
"disableTextWrap": false,
"editorMode": "code",
- "expr": "histogram_quantile(0.99,sum(increase(management_grpc_updatechannel_queue_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
+ "expr": "histogram_quantile(0.99,sum(increase(management_grpc_updatechannel_queue_length_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le,cluster,environment,job))",
"format": "heatmap",
"fullMetaSearch": false,
"hide": false,
@@ -841,7 +841,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "histogram_quantile(0.50, sum(rate(management_store_persistence_duration_ms_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
+ "expr": "histogram_quantile(0.50, sum(rate(management_store_persistence_duration_ms_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
"instant": false,
"legendFormat": "p50",
"range": true,
@@ -853,7 +853,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "histogram_quantile(0.90, sum(rate(management_store_persistence_duration_ms_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
+ "expr": "histogram_quantile(0.90, sum(rate(management_store_persistence_duration_ms_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
"hide": false,
"instant": false,
"legendFormat": "p90",
@@ -866,7 +866,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "histogram_quantile(0.99, sum(rate(management_store_persistence_duration_ms_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
+ "expr": "histogram_quantile(0.99, sum(rate(management_store_persistence_duration_ms_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
"hide": false,
"instant": false,
"legendFormat": "p99",
@@ -963,7 +963,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "histogram_quantile(0.50, sum(rate(management_store_transaction_duration_ms_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
+ "expr": "histogram_quantile(0.50, sum(rate(management_store_transaction_duration_ms_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
"instant": false,
"legendFormat": "p50",
"range": true,
@@ -975,7 +975,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "histogram_quantile(0.90, sum(rate(management_store_transaction_duration_ms_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
+ "expr": "histogram_quantile(0.90, sum(rate(management_store_transaction_duration_ms_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
"hide": false,
"instant": false,
"legendFormat": "p90",
@@ -988,7 +988,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "histogram_quantile(0.99, sum(rate(management_store_transaction_duration_ms_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
+ "expr": "histogram_quantile(0.99, sum(rate(management_store_transaction_duration_ms_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
"hide": false,
"instant": false,
"legendFormat": "p99",
@@ -1085,7 +1085,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "histogram_quantile(0.50, sum(rate(management_store_global_lock_acquisition_duration_ms_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
+ "expr": "histogram_quantile(0.50, sum(rate(management_store_global_lock_acquisition_duration_ms_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
"instant": false,
"legendFormat": "p50",
"range": true,
@@ -1097,7 +1097,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "histogram_quantile(0.90, sum(rate(management_store_global_lock_acquisition_duration_ms_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
+ "expr": "histogram_quantile(0.90, sum(rate(management_store_global_lock_acquisition_duration_ms_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
"hide": false,
"instant": false,
"legendFormat": "p90",
@@ -1110,7 +1110,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "histogram_quantile(0.99, sum(rate(management_store_global_lock_acquisition_duration_ms_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
+ "expr": "histogram_quantile(0.99, sum(rate(management_store_global_lock_acquisition_duration_ms_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (le))",
"hide": false,
"instant": false,
"legendFormat": "p99",
@@ -1221,7 +1221,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "rate(management_idp_authenticate_request_counter_ratio_total{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])",
+ "expr": "rate(management_idp_authenticate_request_counter_total{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])",
"instant": false,
"legendFormat": "{{cluster}}/{{environment}}/{{job}}",
"range": true,
@@ -1317,7 +1317,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "rate(management_idp_get_account_counter_ratio_total{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])",
+ "expr": "rate(management_idp_get_account_counter_total{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])",
"instant": false,
"legendFormat": "{{cluster}}/{{environment}}/{{job}}",
"range": true,
@@ -1413,7 +1413,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "rate(management_idp_update_user_meta_counter_ratio_total{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])",
+ "expr": "rate(management_idp_update_user_meta_counter_total{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])",
"instant": false,
"legendFormat": "{{cluster}}/{{environment}}/{{job}}",
"range": true,
@@ -1523,7 +1523,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "sum(rate(management_http_request_counter_ratio_total{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",method=~\"GET|OPTIONS\"}[$__rate_interval])) by (job,method)",
+ "expr": "sum(rate(management_http_request_counter_total{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",method=~\"GET|OPTIONS\"}[$__rate_interval])) by (job,method)",
"instant": false,
"legendFormat": "{{method}}",
"range": true,
@@ -1619,7 +1619,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "sum(rate(management_http_request_counter_ratio_total{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",method=~\"POST|PUT|DELETE\"}[$__rate_interval])) by (job,method)",
+ "expr": "sum(rate(management_http_request_counter_total{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",method=~\"POST|PUT|DELETE\"}[$__rate_interval])) by (job,method)",
"instant": false,
"legendFormat": "{{method}}",
"range": true,
@@ -1715,7 +1715,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "histogram_quantile(0.50, sum(rate(management_http_request_duration_ms_total_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",type=~\"read\"}[5m])) by (le))",
+ "expr": "histogram_quantile(0.50, sum(rate(management_http_request_duration_ms_total_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",type=~\"read\"}[5m])) by (le))",
"instant": false,
"legendFormat": "p50",
"range": true,
@@ -1727,7 +1727,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "histogram_quantile(0.90, sum(rate(management_http_request_duration_ms_total_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",type=~\"read\"}[5m])) by (le))",
+ "expr": "histogram_quantile(0.90, sum(rate(management_http_request_duration_ms_total_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",type=~\"read\"}[5m])) by (le))",
"hide": false,
"instant": false,
"legendFormat": "p90",
@@ -1740,7 +1740,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "histogram_quantile(0.99, sum(rate(management_http_request_duration_ms_total_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",type=~\"read\"}[5m])) by (le))",
+ "expr": "histogram_quantile(0.99, sum(rate(management_http_request_duration_ms_total_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",type=~\"read\"}[5m])) by (le))",
"hide": false,
"instant": false,
"legendFormat": "p99",
@@ -1837,7 +1837,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "histogram_quantile(0.50, sum(rate(management_http_request_duration_ms_total_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",type=~\"write\"}[5m])) by (le))",
+ "expr": "histogram_quantile(0.50, sum(rate(management_http_request_duration_ms_total_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",type=~\"write\"}[5m])) by (le))",
"instant": false,
"legendFormat": "p50",
"range": true,
@@ -1849,7 +1849,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "histogram_quantile(0.90, sum(rate(management_http_request_duration_ms_total_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",type=~\"write\"}[5m])) by (le))",
+ "expr": "histogram_quantile(0.90, sum(rate(management_http_request_duration_ms_total_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",type=~\"write\"}[5m])) by (le))",
"hide": false,
"instant": false,
"legendFormat": "p90",
@@ -1862,7 +1862,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "histogram_quantile(0.99, sum(rate(management_http_request_duration_ms_total_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",type=~\"write\"}[5m])) by (le))",
+ "expr": "histogram_quantile(0.99, sum(rate(management_http_request_duration_ms_total_milliseconds_bucket{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\",type=~\"write\"}[5m])) by (le))",
"hide": false,
"instant": false,
"legendFormat": "p99",
@@ -1963,7 +1963,7 @@
"uid": "${datasource}"
},
"editorMode": "code",
- "expr": "sum(rate(management_http_request_counter_ratio_total{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (job,exported_endpoint,method)",
+ "expr": "sum(rate(management_http_request_counter_total{cluster=~\"$cluster\",environment=~\"$environment\",job=~\"$job\",host=~\"$host\"}[$__rate_interval])) by (job,exported_endpoint,method)",
"hide": false,
"instant": false,
"legendFormat": "{{method}}-{{exported_endpoint}}",
@@ -3222,7 +3222,7 @@
},
"disableTextWrap": false,
"editorMode": "builder",
- "expr": "sum by(le) (increase(management_grpc_updatechannel_queue_bucket{application=\"management\", environment=\"$environment\", host=~\"$host\"}[$__rate_interval]))",
+ "expr": "sum by(le) (increase(management_grpc_updatechannel_queue_length_bucket{application=\"management\", environment=\"$environment\", host=~\"$host\"}[$__rate_interval]))",
"format": "heatmap",
"fullMetaSearch": false,
"includeNullMetadata": true,
@@ -3323,7 +3323,7 @@
},
"disableTextWrap": false,
"editorMode": "builder",
- "expr": "sum by(le) (increase(management_account_update_account_peers_duration_ms_bucket{application=\"management\", environment=\"$environment\", host=~\"$host\"}[$__rate_interval]))",
+ "expr": "sum by(le) (increase(management_account_update_account_peers_duration_ms_milliseconds_bucket{application=\"management\", environment=\"$environment\", host=~\"$host\"}[$__rate_interval]))",
"format": "heatmap",
"fullMetaSearch": false,
"includeNullMetadata": true,
diff --git a/management/internals/modules/reverseproxy/accesslogs/manager/manager.go b/management/internals/modules/reverseproxy/accesslogs/manager/manager.go
index e8d0ce763..59d7704eb 100644
--- a/management/internals/modules/reverseproxy/accesslogs/manager/manager.go
+++ b/management/internals/modules/reverseproxy/accesslogs/manager/manager.go
@@ -106,13 +106,23 @@ func (m *managerImpl) CleanupOldAccessLogs(ctx context.Context, retentionDays in
// StartPeriodicCleanup starts a background goroutine that periodically cleans up old access logs
func (m *managerImpl) StartPeriodicCleanup(ctx context.Context, retentionDays, cleanupIntervalHours int) {
- if retentionDays <= 0 {
- log.WithContext(ctx).Debug("periodic access log cleanup disabled: retention days is 0 or negative")
+ if retentionDays < 0 {
+ log.WithContext(ctx).Debug("periodic access log cleanup disabled: retention days is negative")
return
}
+ if retentionDays == 0 {
+ retentionDays = 7
+ log.WithContext(ctx).Debugf("no retention days specified for access log cleanup, defaulting to %d days", retentionDays)
+ } else {
+ log.WithContext(ctx).Debugf("access log retention period set to %d days", retentionDays)
+ }
+
if cleanupIntervalHours <= 0 {
cleanupIntervalHours = 24
+ log.WithContext(ctx).Debugf("no cleanup interval specified for access log cleanup, defaulting to %d hours", cleanupIntervalHours)
+ } else {
+ log.WithContext(ctx).Debugf("access log cleanup interval set to %d hours", cleanupIntervalHours)
}
cleanupCtx, cancel := context.WithCancel(ctx)
diff --git a/management/internals/modules/reverseproxy/accesslogs/manager/manager_test.go b/management/internals/modules/reverseproxy/accesslogs/manager/manager_test.go
index 8fadef85f..11bf60829 100644
--- a/management/internals/modules/reverseproxy/accesslogs/manager/manager_test.go
+++ b/management/internals/modules/reverseproxy/accesslogs/manager/manager_test.go
@@ -121,7 +121,7 @@ func TestCleanupWithExactBoundary(t *testing.T) {
}
func TestStartPeriodicCleanup(t *testing.T) {
- t.Run("periodic cleanup disabled with zero retention", func(t *testing.T) {
+ t.Run("periodic cleanup disabled with negative retention", func(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
@@ -135,7 +135,7 @@ func TestStartPeriodicCleanup(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
- manager.StartPeriodicCleanup(ctx, 0, 1)
+ manager.StartPeriodicCleanup(ctx, -1, 1)
time.Sleep(100 * time.Millisecond)
diff --git a/management/internals/modules/reverseproxy/domain/domain.go b/management/internals/modules/reverseproxy/domain/domain.go
index 859f1c5b2..ae13bffae 100644
--- a/management/internals/modules/reverseproxy/domain/domain.go
+++ b/management/internals/modules/reverseproxy/domain/domain.go
@@ -30,3 +30,8 @@ func (d *Domain) EventMeta() map[string]any {
"validated": d.Validated,
}
}
+
+func (d *Domain) Copy() *Domain {
+ dCopy := *d
+ return &dCopy
+}
diff --git a/management/internals/server/config/config.go b/management/internals/server/config/config.go
index 0ba393263..fb9c842b7 100644
--- a/management/internals/server/config/config.go
+++ b/management/internals/server/config/config.go
@@ -203,7 +203,7 @@ type ReverseProxy struct {
// AccessLogRetentionDays specifies the number of days to retain access logs.
// Logs older than this duration will be automatically deleted during cleanup.
- // A value of 0 or negative means logs are kept indefinitely (no cleanup).
+ // A value of 0 will default to 7 days. Negative means logs are kept indefinitely (no cleanup).
AccessLogRetentionDays int
// AccessLogCleanupIntervalHours specifies how often (in hours) to run the cleanup routine.
diff --git a/management/server/account.go b/management/server/account.go
index 75db36a5f..d90b46659 100644
--- a/management/server/account.go
+++ b/management/server/account.go
@@ -742,11 +742,6 @@ func (am *DefaultAccountManager) DeleteAccount(ctx context.Context, accountID, u
return status.Errorf(status.Internal, "failed to build user infos for account %s: %v", accountID, err)
}
- err = am.serviceManager.DeleteAllServices(ctx, accountID, userID)
- if err != nil {
- return status.Errorf(status.Internal, "failed to delete service %s: %v", accountID, err)
- }
-
for _, otherUser := range account.Users {
if otherUser.Id == userID {
continue
diff --git a/management/server/account_test.go b/management/server/account_test.go
index 548cf31d4..2f0533281 100644
--- a/management/server/account_test.go
+++ b/management/server/account_test.go
@@ -15,7 +15,6 @@ import (
"time"
"github.com/golang/mock/gomock"
- "github.com/netbirdio/netbird/shared/management/status"
"github.com/prometheus/client_golang/prometheus/push"
log "github.com/sirupsen/logrus"
"github.com/stretchr/testify/assert"
@@ -23,6 +22,9 @@ import (
"go.opentelemetry.io/otel/metric/noop"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
+ "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
+ "github.com/netbirdio/netbird/shared/management/status"
+
nbdns "github.com/netbirdio/netbird/dns"
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
"github.com/netbirdio/netbird/management/internals/controllers/network_map/controller"
@@ -1815,6 +1817,13 @@ func TestAccount_Copy(t *testing.T) {
Targets: []*service.Target{},
},
},
+ Domains: []*domain.Domain{
+ {
+ ID: "domain1",
+ Domain: "test.com",
+ AccountID: "account1",
+ },
+ },
NetworkMapCache: &types.NetworkMapBuilder{},
}
account.InitOnce()
diff --git a/management/server/migration/migration.go b/management/server/migration/migration.go
index 29555ed0c..7a51cc200 100644
--- a/management/server/migration/migration.go
+++ b/management/server/migration/migration.go
@@ -489,6 +489,102 @@ func MigrateJsonToTable[T any](ctx context.Context, db *gorm.DB, columnName stri
return nil
}
+// hasForeignKey checks whether a foreign key constraint exists on the given table and column.
+func hasForeignKey(db *gorm.DB, table, column string) bool {
+ var count int64
+
+ switch db.Name() {
+ case "postgres":
+ db.Raw(`
+ SELECT COUNT(*) FROM information_schema.key_column_usage kcu
+ JOIN information_schema.table_constraints tc
+ ON tc.constraint_name = kcu.constraint_name
+ AND tc.table_schema = kcu.table_schema
+ WHERE tc.constraint_type = 'FOREIGN KEY'
+ AND kcu.table_name = ?
+ AND kcu.column_name = ?
+ `, table, column).Scan(&count)
+ case "mysql":
+ db.Raw(`
+ SELECT COUNT(*) FROM information_schema.key_column_usage
+ WHERE table_schema = DATABASE()
+ AND table_name = ?
+ AND column_name = ?
+ AND referenced_table_name IS NOT NULL
+ `, table, column).Scan(&count)
+ default: // sqlite
+ type fkInfo struct {
+ From string
+ }
+ var fks []fkInfo
+ db.Raw(fmt.Sprintf("PRAGMA foreign_key_list(%s)", table)).Scan(&fks)
+ for _, fk := range fks {
+ if fk.From == column {
+ return true
+ }
+ }
+ return false
+ }
+
+ return count > 0
+}
+
+// CleanupOrphanedResources deletes rows from the table of model T where the foreign
+// key column (fkColumn) references a row in the table of model R that no longer exists.
+func CleanupOrphanedResources[T any, R any](ctx context.Context, db *gorm.DB, fkColumn string) error {
+ var model T
+ var refModel R
+
+ if !db.Migrator().HasTable(&model) {
+ log.WithContext(ctx).Debugf("table for %T does not exist, no cleanup needed", model)
+ return nil
+ }
+
+ if !db.Migrator().HasTable(&refModel) {
+ log.WithContext(ctx).Debugf("referenced table for %T does not exist, no cleanup needed", refModel)
+ return nil
+ }
+
+ stmtT := &gorm.Statement{DB: db}
+ if err := stmtT.Parse(&model); err != nil {
+ return fmt.Errorf("parse model %T: %w", model, err)
+ }
+ childTable := stmtT.Schema.Table
+
+ stmtR := &gorm.Statement{DB: db}
+ if err := stmtR.Parse(&refModel); err != nil {
+ return fmt.Errorf("parse reference model %T: %w", refModel, err)
+ }
+ parentTable := stmtR.Schema.Table
+
+ if !db.Migrator().HasColumn(&model, fkColumn) {
+ log.WithContext(ctx).Debugf("column %s does not exist in table %s, no cleanup needed", fkColumn, childTable)
+ return nil
+ }
+
+ // If a foreign key constraint already exists on the column, the DB itself
+ // enforces referential integrity and orphaned rows cannot exist.
+ if hasForeignKey(db, childTable, fkColumn) {
+ log.WithContext(ctx).Debugf("foreign key constraint for %s already exists on %s, no cleanup needed", fkColumn, childTable)
+ return nil
+ }
+
+ result := db.Exec(
+ fmt.Sprintf(
+ "DELETE FROM %s WHERE %s NOT IN (SELECT id FROM %s)",
+ childTable, fkColumn, parentTable,
+ ),
+ )
+ if result.Error != nil {
+ return fmt.Errorf("cleanup orphaned rows in %s: %w", childTable, result.Error)
+ }
+
+ log.WithContext(ctx).Infof("Cleaned up %d orphaned rows from %s where %s had no matching row in %s",
+ result.RowsAffected, childTable, fkColumn, parentTable)
+
+ return nil
+}
+
func RemoveDuplicatePeerKeys(ctx context.Context, db *gorm.DB) error {
if !db.Migrator().HasTable("peers") {
log.WithContext(ctx).Debug("peers table does not exist, skipping duplicate key cleanup")
diff --git a/management/server/migration/migration_test.go b/management/server/migration/migration_test.go
index c1be8a3a3..5e00976c2 100644
--- a/management/server/migration/migration_test.go
+++ b/management/server/migration/migration_test.go
@@ -441,3 +441,197 @@ func TestRemoveDuplicatePeerKeys_NoTable(t *testing.T) {
err := migration.RemoveDuplicatePeerKeys(context.Background(), db)
require.NoError(t, err, "Should not fail when table does not exist")
}
+
+type testParent struct {
+ ID string `gorm:"primaryKey"`
+}
+
+func (testParent) TableName() string {
+ return "test_parents"
+}
+
+type testChild struct {
+ ID string `gorm:"primaryKey"`
+ ParentID string
+}
+
+func (testChild) TableName() string {
+ return "test_children"
+}
+
+type testChildWithFK struct {
+ ID string `gorm:"primaryKey"`
+ ParentID string `gorm:"index"`
+ Parent *testParent `gorm:"foreignKey:ParentID"`
+}
+
+func (testChildWithFK) TableName() string {
+ return "test_children"
+}
+
+func setupOrphanTestDB(t *testing.T, models ...any) *gorm.DB {
+ t.Helper()
+ db := setupDatabase(t)
+ for _, m := range models {
+ _ = db.Migrator().DropTable(m)
+ }
+ err := db.AutoMigrate(models...)
+ require.NoError(t, err, "Failed to auto-migrate tables")
+ return db
+}
+
+func TestCleanupOrphanedResources_NoChildTable(t *testing.T) {
+ db := setupDatabase(t)
+ _ = db.Migrator().DropTable(&testChild{})
+ _ = db.Migrator().DropTable(&testParent{})
+
+ err := migration.CleanupOrphanedResources[testChild, testParent](context.Background(), db, "parent_id")
+ require.NoError(t, err, "Should not fail when child table does not exist")
+}
+
+func TestCleanupOrphanedResources_NoParentTable(t *testing.T) {
+ db := setupDatabase(t)
+ _ = db.Migrator().DropTable(&testParent{})
+ _ = db.Migrator().DropTable(&testChild{})
+
+ err := db.AutoMigrate(&testChild{})
+ require.NoError(t, err)
+
+ err = migration.CleanupOrphanedResources[testChild, testParent](context.Background(), db, "parent_id")
+ require.NoError(t, err, "Should not fail when parent table does not exist")
+}
+
+func TestCleanupOrphanedResources_EmptyTables(t *testing.T) {
+ db := setupOrphanTestDB(t, &testParent{}, &testChild{})
+
+ err := migration.CleanupOrphanedResources[testChild, testParent](context.Background(), db, "parent_id")
+ require.NoError(t, err, "Should not fail on empty tables")
+
+ var count int64
+ db.Model(&testChild{}).Count(&count)
+ assert.Equal(t, int64(0), count)
+}
+
+func TestCleanupOrphanedResources_NoOrphans(t *testing.T) {
+ db := setupOrphanTestDB(t, &testParent{}, &testChild{})
+
+ require.NoError(t, db.Create(&testParent{ID: "p1"}).Error)
+ require.NoError(t, db.Create(&testParent{ID: "p2"}).Error)
+ require.NoError(t, db.Create(&testChild{ID: "c1", ParentID: "p1"}).Error)
+ require.NoError(t, db.Create(&testChild{ID: "c2", ParentID: "p2"}).Error)
+
+ err := migration.CleanupOrphanedResources[testChild, testParent](context.Background(), db, "parent_id")
+ require.NoError(t, err)
+
+ var count int64
+ db.Model(&testChild{}).Count(&count)
+ assert.Equal(t, int64(2), count, "All children should remain when no orphans")
+}
+
+func TestCleanupOrphanedResources_AllOrphans(t *testing.T) {
+ db := setupOrphanTestDB(t, &testParent{}, &testChild{})
+
+ require.NoError(t, db.Exec("INSERT INTO test_children (id, parent_id) VALUES (?, ?)", "c1", "gone1").Error)
+ require.NoError(t, db.Exec("INSERT INTO test_children (id, parent_id) VALUES (?, ?)", "c2", "gone2").Error)
+ require.NoError(t, db.Exec("INSERT INTO test_children (id, parent_id) VALUES (?, ?)", "c3", "gone3").Error)
+
+ err := migration.CleanupOrphanedResources[testChild, testParent](context.Background(), db, "parent_id")
+ require.NoError(t, err)
+
+ var count int64
+ db.Model(&testChild{}).Count(&count)
+ assert.Equal(t, int64(0), count, "All orphaned children should be deleted")
+}
+
+func TestCleanupOrphanedResources_MixedValidAndOrphaned(t *testing.T) {
+ db := setupOrphanTestDB(t, &testParent{}, &testChild{})
+
+ require.NoError(t, db.Create(&testParent{ID: "p1"}).Error)
+ require.NoError(t, db.Create(&testParent{ID: "p2"}).Error)
+
+ require.NoError(t, db.Create(&testChild{ID: "c1", ParentID: "p1"}).Error)
+ require.NoError(t, db.Create(&testChild{ID: "c2", ParentID: "p2"}).Error)
+ require.NoError(t, db.Create(&testChild{ID: "c3", ParentID: "p1"}).Error)
+
+ require.NoError(t, db.Exec("INSERT INTO test_children (id, parent_id) VALUES (?, ?)", "c4", "gone1").Error)
+ require.NoError(t, db.Exec("INSERT INTO test_children (id, parent_id) VALUES (?, ?)", "c5", "gone2").Error)
+
+ err := migration.CleanupOrphanedResources[testChild, testParent](context.Background(), db, "parent_id")
+ require.NoError(t, err)
+
+ var remaining []testChild
+ require.NoError(t, db.Order("id").Find(&remaining).Error)
+
+ assert.Len(t, remaining, 3, "Only valid children should remain")
+ assert.Equal(t, "c1", remaining[0].ID)
+ assert.Equal(t, "c2", remaining[1].ID)
+ assert.Equal(t, "c3", remaining[2].ID)
+}
+
+func TestCleanupOrphanedResources_Idempotent(t *testing.T) {
+ db := setupOrphanTestDB(t, &testParent{}, &testChild{})
+
+ require.NoError(t, db.Create(&testParent{ID: "p1"}).Error)
+ require.NoError(t, db.Create(&testChild{ID: "c1", ParentID: "p1"}).Error)
+ require.NoError(t, db.Exec("INSERT INTO test_children (id, parent_id) VALUES (?, ?)", "c2", "gone").Error)
+
+ ctx := context.Background()
+
+ err := migration.CleanupOrphanedResources[testChild, testParent](ctx, db, "parent_id")
+ require.NoError(t, err)
+
+ var count int64
+ db.Model(&testChild{}).Count(&count)
+ assert.Equal(t, int64(1), count)
+
+ err = migration.CleanupOrphanedResources[testChild, testParent](ctx, db, "parent_id")
+ require.NoError(t, err)
+
+ db.Model(&testChild{}).Count(&count)
+ assert.Equal(t, int64(1), count, "Count should remain the same after second run")
+}
+
+func TestCleanupOrphanedResources_SkipsWhenForeignKeyExists(t *testing.T) {
+ engine := os.Getenv("NETBIRD_STORE_ENGINE")
+ if engine != "postgres" && engine != "mysql" {
+ t.Skip("FK constraint early-exit test requires postgres or mysql")
+ }
+
+ db := setupDatabase(t)
+ _ = db.Migrator().DropTable(&testChildWithFK{})
+ _ = db.Migrator().DropTable(&testParent{})
+
+ err := db.AutoMigrate(&testParent{}, &testChildWithFK{})
+ require.NoError(t, err)
+
+ require.NoError(t, db.Create(&testParent{ID: "p1"}).Error)
+ require.NoError(t, db.Create(&testParent{ID: "p2"}).Error)
+ require.NoError(t, db.Create(&testChildWithFK{ID: "c1", ParentID: "p1"}).Error)
+ require.NoError(t, db.Create(&testChildWithFK{ID: "c2", ParentID: "p2"}).Error)
+
+ switch engine {
+ case "postgres":
+ require.NoError(t, db.Exec("ALTER TABLE test_children DROP CONSTRAINT fk_test_children_parent").Error)
+ require.NoError(t, db.Exec("DELETE FROM test_parents WHERE id = ?", "p2").Error)
+ require.NoError(t, db.Exec(
+ "ALTER TABLE test_children ADD CONSTRAINT fk_test_children_parent "+
+ "FOREIGN KEY (parent_id) REFERENCES test_parents(id) NOT VALID",
+ ).Error)
+ case "mysql":
+ require.NoError(t, db.Exec("SET FOREIGN_KEY_CHECKS = 0").Error)
+ require.NoError(t, db.Exec("ALTER TABLE test_children DROP FOREIGN KEY fk_test_children_parent").Error)
+ require.NoError(t, db.Exec("DELETE FROM test_parents WHERE id = ?", "p2").Error)
+ require.NoError(t, db.Exec(
+ "ALTER TABLE test_children ADD CONSTRAINT fk_test_children_parent "+
+ "FOREIGN KEY (parent_id) REFERENCES test_parents(id)",
+ ).Error)
+ require.NoError(t, db.Exec("SET FOREIGN_KEY_CHECKS = 1").Error)
+ }
+
+ err = migration.CleanupOrphanedResources[testChildWithFK, testParent](context.Background(), db, "parent_id")
+ require.NoError(t, err)
+
+ var count int64
+ db.Model(&testChildWithFK{}).Count(&count)
+ assert.Equal(t, int64(2), count, "Both rows should survive — migration must skip when FK constraint exists")
+}
diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go
index 397b8673d..0b463a724 100644
--- a/management/server/store/sql_store.go
+++ b/management/server/store/sql_store.go
@@ -396,6 +396,11 @@ func (s *SqlStore) DeleteAccount(ctx context.Context, account *types.Account) er
return result.Error
}
+ result = tx.Select(clause.Associations).Delete(account.Services, "account_id = ?", account.Id)
+ if result.Error != nil {
+ return result.Error
+ }
+
result = tx.Select(clause.Associations).Delete(account)
if result.Error != nil {
return result.Error
@@ -2099,6 +2104,8 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv
var createdAt, certIssuedAt sql.NullTime
var status, proxyCluster, sessionPrivateKey, sessionPublicKey sql.NullString
var mode, source, sourcePeer sql.NullString
+ var terminated, portAutoAssigned sql.NullBool
+ var listenPort sql.NullInt64
err := row.Scan(
&s.ID,
&s.AccountID,
@@ -2115,11 +2122,11 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv
&sessionPrivateKey,
&sessionPublicKey,
&mode,
- &s.ListenPort,
- &s.PortAutoAssigned,
+ &listenPort,
+ &portAutoAssigned,
&source,
&sourcePeer,
- &s.Terminated,
+ &terminated,
)
if err != nil {
return nil, err
@@ -2160,7 +2167,15 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv
if sourcePeer.Valid {
s.SourcePeer = sourcePeer.String
}
-
+ if terminated.Valid {
+ s.Terminated = terminated.Bool
+ }
+ if portAutoAssigned.Valid {
+ s.PortAutoAssigned = portAutoAssigned.Bool
+ }
+ if listenPort.Valid {
+ s.ListenPort = uint16(listenPort.Int64)
+ }
s.Targets = []*rpservice.Target{}
return &s, nil
})
diff --git a/management/server/store/sql_store_test.go b/management/server/store/sql_store_test.go
index bafa63580..8ea6c2ae5 100644
--- a/management/server/store/sql_store_test.go
+++ b/management/server/store/sql_store_test.go
@@ -22,6 +22,8 @@ import (
"github.com/stretchr/testify/require"
nbdns "github.com/netbirdio/netbird/dns"
+ proxydomain "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
+ rpservice "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"
@@ -350,6 +352,35 @@ func TestSqlite_DeleteAccount(t *testing.T) {
},
}
+ account.Services = []*rpservice.Service{
+ {
+ ID: "service_id",
+ AccountID: account.Id,
+ Name: "test service",
+ Domain: "svc.example.com",
+ Enabled: true,
+ Targets: []*rpservice.Target{
+ {
+ AccountID: account.Id,
+ ServiceID: "service_id",
+ Host: "localhost",
+ Port: 8080,
+ Protocol: "http",
+ Enabled: true,
+ },
+ },
+ },
+ }
+
+ account.Domains = []*proxydomain.Domain{
+ {
+ ID: "domain_id",
+ Domain: "custom.example.com",
+ AccountID: account.Id,
+ Validated: true,
+ },
+ }
+
err = store.SaveAccount(context.Background(), account)
require.NoError(t, err)
@@ -411,6 +442,20 @@ func TestSqlite_DeleteAccount(t *testing.T) {
require.NoError(t, err, "expecting no error after removing DeleteAccount when searching for network resources")
require.Len(t, resources, 0, "expecting no network resources to be found after DeleteAccount")
}
+
+ domains, err := store.ListCustomDomains(context.Background(), account.Id)
+ require.NoError(t, err, "expecting no error after DeleteAccount when searching for custom domains")
+ require.Len(t, domains, 0, "expecting no custom domains to be found after DeleteAccount")
+
+ var services []*rpservice.Service
+ err = store.(*SqlStore).db.Model(&rpservice.Service{}).Find(&services, "account_id = ?", account.Id).Error
+ require.NoError(t, err, "expecting no error after DeleteAccount when searching for services")
+ require.Len(t, services, 0, "expecting no services to be found after DeleteAccount")
+
+ var targets []*rpservice.Target
+ err = store.(*SqlStore).db.Model(&rpservice.Target{}).Find(&targets, "account_id = ?", account.Id).Error
+ require.NoError(t, err, "expecting no error after DeleteAccount when searching for service targets")
+ require.Len(t, targets, 0, "expecting no service targets to be found after DeleteAccount")
}
func Test_GetAccount(t *testing.T) {
diff --git a/management/server/store/sqlstore_bench_test.go b/management/server/store/sqlstore_bench_test.go
index f2abafceb..81c4b33ae 100644
--- a/management/server/store/sqlstore_bench_test.go
+++ b/management/server/store/sqlstore_bench_test.go
@@ -20,6 +20,7 @@ import (
"github.com/stretchr/testify/assert"
nbdns "github.com/netbirdio/netbird/dns"
+ "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
@@ -265,6 +266,7 @@ func setupBenchmarkDB(b testing.TB) (*SqlStore, func(), string) {
&nbdns.NameServerGroup{}, &posture.Checks{}, &networkTypes.Network{},
&routerTypes.NetworkRouter{}, &resourceTypes.NetworkResource{},
&types.AccountOnboarding{}, &service.Service{}, &service.Target{},
+ &domain.Domain{},
}
for i := len(models) - 1; i >= 0; i-- {
diff --git a/management/server/store/store.go b/management/server/store/store.go
index f0c34ffa9..efd9a28fd 100644
--- a/management/server/store/store.go
+++ b/management/server/store/store.go
@@ -448,6 +448,12 @@ func getMigrationsPreAuto(ctx context.Context) []migrationFunc {
func(db *gorm.DB) error {
return migration.RemoveDuplicatePeerKeys(ctx, db)
},
+ func(db *gorm.DB) error {
+ return migration.CleanupOrphanedResources[rpservice.Service, types.Account](ctx, db, "account_id")
+ },
+ func(db *gorm.DB) error {
+ return migration.CleanupOrphanedResources[domain.Domain, types.Account](ctx, db, "account_id")
+ },
}
}
diff --git a/management/server/types/account.go b/management/server/types/account.go
index 269fc7a88..c448813db 100644
--- a/management/server/types/account.go
+++ b/management/server/types/account.go
@@ -18,6 +18,7 @@ import (
"github.com/netbirdio/netbird/client/ssh/auth"
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"
@@ -101,6 +102,7 @@ type Account struct {
DNSSettings DNSSettings `gorm:"embedded;embeddedPrefix:dns_settings_"`
PostureChecks []*posture.Checks `gorm:"foreignKey:AccountID;references:id"`
Services []*service.Service `gorm:"foreignKey:AccountID;references:id"`
+ Domains []*proxydomain.Domain `gorm:"foreignKey:AccountID;references:id"`
// Settings is a dictionary of Account settings
Settings *Settings `gorm:"embedded;embeddedPrefix:settings_"`
Networks []*networkTypes.Network `gorm:"foreignKey:AccountID;references:id"`
@@ -911,6 +913,11 @@ func (a *Account) Copy() *Account {
services = append(services, svc.Copy())
}
+ domains := []*proxydomain.Domain{}
+ for _, domain := range a.Domains {
+ domains = append(domains, domain.Copy())
+ }
+
return &Account{
Id: a.Id,
CreatedBy: a.CreatedBy,
@@ -936,6 +943,7 @@ func (a *Account) Copy() *Account {
Onboarding: a.Onboarding,
NetworkMapCache: a.NetworkMapCache,
nmapInitOnce: a.nmapInitOnce,
+ Domains: domains,
}
}
diff --git a/management/server/types/networkmap_benchmark_test.go b/management/server/types/networkmap_benchmark_test.go
new file mode 100644
index 000000000..38272e7b0
--- /dev/null
+++ b/management/server/types/networkmap_benchmark_test.go
@@ -0,0 +1,217 @@
+package types_test
+
+import (
+ "context"
+ "fmt"
+ "os"
+ "testing"
+
+ nbdns "github.com/netbirdio/netbird/dns"
+ "github.com/netbirdio/netbird/management/server/types"
+)
+
+type benchmarkScale struct {
+ name string
+ peers int
+ groups int
+}
+
+var defaultScales = []benchmarkScale{
+ {"100peers_5groups", 100, 5},
+ {"500peers_20groups", 500, 20},
+ {"1000peers_50groups", 1000, 50},
+ {"5000peers_100groups", 5000, 100},
+ {"10000peers_200groups", 10000, 200},
+ {"20000peers_200groups", 20000, 200},
+ {"30000peers_300groups", 30000, 300},
+}
+
+func skipCIBenchmark(b *testing.B) {
+ if os.Getenv("CI") == "true" {
+ b.Skip("Skipping benchmark in CI")
+ }
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// Single Peer Network Map Generation
+// ──────────────────────────────────────────────────────────────────────────────
+
+// BenchmarkNetworkMapGeneration_Components benchmarks the components-based approach for a single peer.
+func BenchmarkNetworkMapGeneration_Components(b *testing.B) {
+ skipCIBenchmark(b)
+ for _, scale := range defaultScales {
+ b.Run(scale.name, func(b *testing.B) {
+ account, validatedPeers := scalableTestAccount(scale.peers, scale.groups)
+ ctx := context.Background()
+ resourcePolicies := account.GetResourcePoliciesMap()
+ routers := account.GetResourceRoutersMap()
+ groupIDToUserIDs := account.GetActiveGroupUsers()
+
+ b.ReportAllocs()
+ b.ResetTimer()
+ for range b.N {
+ _ = account.GetPeerNetworkMapFromComponents(ctx, "peer-0", nbdns.CustomZone{}, nil, validatedPeers, resourcePolicies, routers, nil, groupIDToUserIDs)
+ }
+ })
+ }
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// All Peers (UpdateAccountPeers hot path)
+// ──────────────────────────────────────────────────────────────────────────────
+
+// BenchmarkNetworkMapGeneration_AllPeers benchmarks generating network maps for ALL peers.
+func BenchmarkNetworkMapGeneration_AllPeers(b *testing.B) {
+ skipCIBenchmark(b)
+ scales := []benchmarkScale{
+ {"100peers_5groups", 100, 5},
+ {"500peers_20groups", 500, 20},
+ {"1000peers_50groups", 1000, 50},
+ {"5000peers_100groups", 5000, 100},
+ }
+
+ for _, scale := range scales {
+ account, validatedPeers := scalableTestAccount(scale.peers, scale.groups)
+ ctx := context.Background()
+
+ peerIDs := make([]string, 0, len(account.Peers))
+ for peerID := range account.Peers {
+ peerIDs = append(peerIDs, peerID)
+ }
+
+ b.Run("components/"+scale.name, func(b *testing.B) {
+ resourcePolicies := account.GetResourcePoliciesMap()
+ routers := account.GetResourceRoutersMap()
+ groupIDToUserIDs := account.GetActiveGroupUsers()
+ b.ReportAllocs()
+ b.ResetTimer()
+ for range b.N {
+ for _, peerID := range peerIDs {
+ _ = account.GetPeerNetworkMapFromComponents(ctx, peerID, nbdns.CustomZone{}, nil, validatedPeers, resourcePolicies, routers, nil, groupIDToUserIDs)
+ }
+ }
+ })
+ }
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// Sub-operations
+// ──────────────────────────────────────────────────────────────────────────────
+
+// BenchmarkNetworkMapGeneration_ComponentsCreation benchmarks components extraction.
+func BenchmarkNetworkMapGeneration_ComponentsCreation(b *testing.B) {
+ skipCIBenchmark(b)
+ for _, scale := range defaultScales {
+ b.Run(scale.name, func(b *testing.B) {
+ account, validatedPeers := scalableTestAccount(scale.peers, scale.groups)
+ ctx := context.Background()
+ resourcePolicies := account.GetResourcePoliciesMap()
+ routers := account.GetResourceRoutersMap()
+ groupIDToUserIDs := account.GetActiveGroupUsers()
+ b.ReportAllocs()
+ b.ResetTimer()
+ for range b.N {
+ _ = account.GetPeerNetworkMapComponents(ctx, "peer-0", nbdns.CustomZone{}, nil, validatedPeers, resourcePolicies, routers, groupIDToUserIDs)
+ }
+ })
+ }
+}
+
+// BenchmarkNetworkMapGeneration_ComponentsCalculation benchmarks calculation from pre-built components.
+func BenchmarkNetworkMapGeneration_ComponentsCalculation(b *testing.B) {
+ skipCIBenchmark(b)
+ for _, scale := range defaultScales {
+ b.Run(scale.name, func(b *testing.B) {
+ account, validatedPeers := scalableTestAccount(scale.peers, scale.groups)
+ ctx := context.Background()
+ resourcePolicies := account.GetResourcePoliciesMap()
+ routers := account.GetResourceRoutersMap()
+ groupIDToUserIDs := account.GetActiveGroupUsers()
+ components := account.GetPeerNetworkMapComponents(ctx, "peer-0", nbdns.CustomZone{}, nil, validatedPeers, resourcePolicies, routers, groupIDToUserIDs)
+ b.ReportAllocs()
+ b.ResetTimer()
+ for range b.N {
+ _ = types.CalculateNetworkMapFromComponents(ctx, components)
+ }
+ })
+ }
+}
+
+// BenchmarkNetworkMapGeneration_PrecomputeMaps benchmarks precomputed map costs.
+func BenchmarkNetworkMapGeneration_PrecomputeMaps(b *testing.B) {
+ skipCIBenchmark(b)
+ for _, scale := range defaultScales {
+ b.Run("ResourcePoliciesMap/"+scale.name, func(b *testing.B) {
+ account, _ := scalableTestAccount(scale.peers, scale.groups)
+ b.ReportAllocs()
+ b.ResetTimer()
+ for range b.N {
+ _ = account.GetResourcePoliciesMap()
+ }
+ })
+ b.Run("ResourceRoutersMap/"+scale.name, func(b *testing.B) {
+ account, _ := scalableTestAccount(scale.peers, scale.groups)
+ b.ReportAllocs()
+ b.ResetTimer()
+ for range b.N {
+ _ = account.GetResourceRoutersMap()
+ }
+ })
+ b.Run("ActiveGroupUsers/"+scale.name, func(b *testing.B) {
+ account, _ := scalableTestAccount(scale.peers, scale.groups)
+ b.ReportAllocs()
+ b.ResetTimer()
+ for range b.N {
+ _ = account.GetActiveGroupUsers()
+ }
+ })
+ }
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// Scaling Analysis
+// ──────────────────────────────────────────────────────────────────────────────
+
+// BenchmarkNetworkMapGeneration_GroupScaling tests group count impact on performance.
+func BenchmarkNetworkMapGeneration_GroupScaling(b *testing.B) {
+ skipCIBenchmark(b)
+ groupCounts := []int{1, 5, 20, 50, 100, 200, 500}
+ for _, numGroups := range groupCounts {
+ b.Run(fmt.Sprintf("components_%dgroups", numGroups), func(b *testing.B) {
+ account, validatedPeers := scalableTestAccount(1000, numGroups)
+ ctx := context.Background()
+ resourcePolicies := account.GetResourcePoliciesMap()
+ routers := account.GetResourceRoutersMap()
+ groupIDToUserIDs := account.GetActiveGroupUsers()
+ b.ReportAllocs()
+ b.ResetTimer()
+ for range b.N {
+ _ = account.GetPeerNetworkMapFromComponents(ctx, "peer-0", nbdns.CustomZone{}, nil, validatedPeers, resourcePolicies, routers, nil, groupIDToUserIDs)
+ }
+ })
+ }
+}
+
+// BenchmarkNetworkMapGeneration_PeerScaling tests peer count impact on performance.
+func BenchmarkNetworkMapGeneration_PeerScaling(b *testing.B) {
+ skipCIBenchmark(b)
+ peerCounts := []int{50, 100, 500, 1000, 2000, 5000, 10000, 20000, 30000}
+ for _, numPeers := range peerCounts {
+ numGroups := numPeers / 20
+ if numGroups < 1 {
+ numGroups = 1
+ }
+ b.Run(fmt.Sprintf("components_%dpeers", numPeers), func(b *testing.B) {
+ account, validatedPeers := scalableTestAccount(numPeers, numGroups)
+ ctx := context.Background()
+ resourcePolicies := account.GetResourcePoliciesMap()
+ routers := account.GetResourceRoutersMap()
+ groupIDToUserIDs := account.GetActiveGroupUsers()
+ b.ReportAllocs()
+ b.ResetTimer()
+ for range b.N {
+ _ = account.GetPeerNetworkMapFromComponents(ctx, "peer-0", nbdns.CustomZone{}, nil, validatedPeers, resourcePolicies, routers, nil, groupIDToUserIDs)
+ }
+ })
+ }
+}
diff --git a/management/server/types/networkmap_components_correctness_test.go b/management/server/types/networkmap_components_correctness_test.go
new file mode 100644
index 000000000..5cd41ff10
--- /dev/null
+++ b/management/server/types/networkmap_components_correctness_test.go
@@ -0,0 +1,1192 @@
+package types_test
+
+import (
+ "context"
+ "fmt"
+ "net"
+ "net/netip"
+ "testing"
+ "time"
+
+ "github.com/stretchr/testify/assert"
+ "github.com/stretchr/testify/require"
+
+ 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"
+ networkTypes "github.com/netbirdio/netbird/management/server/networks/types"
+ 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/route"
+)
+
+// scalableTestAccountWithoutDefaultPolicy creates an account without the blanket "Allow All" policy.
+// Use this for tests that need to verify feature-specific connectivity in isolation.
+func scalableTestAccountWithoutDefaultPolicy(numPeers, numGroups int) (*types.Account, map[string]struct{}) {
+ return buildScalableTestAccount(numPeers, numGroups, false)
+}
+
+// scalableTestAccount creates a realistic account with a blanket "Allow All" policy
+// plus per-group policies, routes, network resources, posture checks, and DNS settings.
+func scalableTestAccount(numPeers, numGroups int) (*types.Account, map[string]struct{}) {
+ return buildScalableTestAccount(numPeers, numGroups, true)
+}
+
+// buildScalableTestAccount is the core builder. When withDefaultPolicy is true it adds
+// a blanket group-all <-> group-all allow rule; when false the only policies are the
+// per-group ones, so tests can verify feature-specific connectivity in isolation.
+func buildScalableTestAccount(numPeers, numGroups int, withDefaultPolicy bool) (*types.Account, map[string]struct{}) {
+ peers := make(map[string]*nbpeer.Peer, numPeers)
+ allGroupPeers := make([]string, 0, numPeers)
+
+ for i := range numPeers {
+ peerID := fmt.Sprintf("peer-%d", i)
+ ip := net.IP{100, byte(64 + i/65536), byte((i / 256) % 256), byte(i % 256)}
+ wtVersion := "0.25.0"
+ if i%2 == 0 {
+ wtVersion = "0.40.0"
+ }
+
+ p := &nbpeer.Peer{
+ ID: peerID,
+ IP: ip,
+ Key: fmt.Sprintf("key-%s", peerID),
+ DNSLabel: fmt.Sprintf("peer%d", i),
+ Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now()},
+ UserID: "user-admin",
+ Meta: nbpeer.PeerSystemMeta{WtVersion: wtVersion, GoOS: "linux"},
+ }
+
+ if i == numPeers-2 {
+ p.LoginExpirationEnabled = true
+ pastTimestamp := time.Now().Add(-2 * time.Hour)
+ p.LastLogin = &pastTimestamp
+ }
+
+ peers[peerID] = p
+ allGroupPeers = append(allGroupPeers, peerID)
+ }
+
+ groups := make(map[string]*types.Group, numGroups+1)
+ groups["group-all"] = &types.Group{ID: "group-all", Name: "All", Peers: allGroupPeers}
+
+ peersPerGroup := numPeers / numGroups
+ if peersPerGroup < 1 {
+ peersPerGroup = 1
+ }
+
+ for g := range numGroups {
+ groupID := fmt.Sprintf("group-%d", g)
+ groupPeers := make([]string, 0, peersPerGroup)
+ start := g * peersPerGroup
+ end := start + peersPerGroup
+ if end > numPeers {
+ end = numPeers
+ }
+ for i := start; i < end; i++ {
+ groupPeers = append(groupPeers, fmt.Sprintf("peer-%d", i))
+ }
+ groups[groupID] = &types.Group{ID: groupID, Name: fmt.Sprintf("Group %d", g), Peers: groupPeers}
+ }
+
+ policies := make([]*types.Policy, 0, numGroups+2)
+ if withDefaultPolicy {
+ policies = append(policies, &types.Policy{
+ ID: "policy-all", Name: "Default-Allow", Enabled: true,
+ Rules: []*types.PolicyRule{{
+ ID: "rule-all", Name: "Allow All", Enabled: true, Action: types.PolicyTrafficActionAccept,
+ Protocol: types.PolicyRuleProtocolALL, Bidirectional: true,
+ Sources: []string{"group-all"}, Destinations: []string{"group-all"},
+ }},
+ })
+ }
+
+ for g := range numGroups {
+ groupID := fmt.Sprintf("group-%d", g)
+ dstGroup := fmt.Sprintf("group-%d", (g+1)%numGroups)
+ policies = append(policies, &types.Policy{
+ ID: fmt.Sprintf("policy-%d", g), Name: fmt.Sprintf("Policy %d", g), Enabled: true,
+ Rules: []*types.PolicyRule{{
+ ID: fmt.Sprintf("rule-%d", g), Name: fmt.Sprintf("Rule %d", g), Enabled: true,
+ Action: types.PolicyTrafficActionAccept, Protocol: types.PolicyRuleProtocolTCP,
+ Bidirectional: true,
+ Ports: []string{"8080"},
+ Sources: []string{groupID}, Destinations: []string{dstGroup},
+ }},
+ })
+ }
+
+ if numGroups >= 2 {
+ policies = append(policies, &types.Policy{
+ ID: "policy-drop", Name: "Drop DB traffic", Enabled: true,
+ Rules: []*types.PolicyRule{{
+ ID: "rule-drop", Name: "Drop DB", Enabled: true, Action: types.PolicyTrafficActionDrop,
+ Protocol: types.PolicyRuleProtocolTCP, Ports: []string{"5432"}, Bidirectional: true,
+ Sources: []string{"group-0"}, Destinations: []string{"group-1"},
+ }},
+ })
+ }
+
+ numRoutes := numGroups
+ if numRoutes > 20 {
+ numRoutes = 20
+ }
+ routes := make(map[route.ID]*route.Route, numRoutes)
+ for r := range numRoutes {
+ routeID := route.ID(fmt.Sprintf("route-%d", r))
+ peerIdx := (numPeers / 2) + r
+ if peerIdx >= numPeers {
+ peerIdx = numPeers - 1
+ }
+ routePeerID := fmt.Sprintf("peer-%d", peerIdx)
+ groupID := fmt.Sprintf("group-%d", r%numGroups)
+ routes[routeID] = &route.Route{
+ ID: routeID,
+ Network: netip.MustParsePrefix(fmt.Sprintf("10.%d.0.0/16", r)),
+ Peer: peers[routePeerID].Key,
+ PeerID: routePeerID,
+ Description: fmt.Sprintf("Route %d", r),
+ Enabled: true,
+ PeerGroups: []string{groupID},
+ Groups: []string{"group-all"},
+ AccessControlGroups: []string{groupID},
+ AccountID: "test-account",
+ }
+ }
+
+ numResources := numGroups / 2
+ if numResources < 1 {
+ numResources = 1
+ }
+ if numResources > 50 {
+ numResources = 50
+ }
+
+ networkResources := make([]*resourceTypes.NetworkResource, 0, numResources)
+ networksList := make([]*networkTypes.Network, 0, numResources)
+ networkRouters := make([]*routerTypes.NetworkRouter, 0, numResources)
+
+ routingPeerStart := numPeers * 3 / 4
+ for nr := range numResources {
+ netID := fmt.Sprintf("net-%d", nr)
+ resID := fmt.Sprintf("res-%d", nr)
+ routerPeerIdx := routingPeerStart + nr
+ if routerPeerIdx >= numPeers {
+ routerPeerIdx = numPeers - 1
+ }
+ routerPeerID := fmt.Sprintf("peer-%d", routerPeerIdx)
+
+ networksList = append(networksList, &networkTypes.Network{ID: netID, Name: fmt.Sprintf("Network %d", nr), AccountID: "test-account"})
+ networkResources = append(networkResources, &resourceTypes.NetworkResource{
+ ID: resID, NetworkID: netID, AccountID: "test-account", Enabled: true,
+ Address: fmt.Sprintf("svc-%d.netbird.cloud", nr),
+ })
+ networkRouters = append(networkRouters, &routerTypes.NetworkRouter{
+ ID: fmt.Sprintf("router-%d", nr), NetworkID: netID, Peer: routerPeerID,
+ Enabled: true, AccountID: "test-account",
+ })
+
+ policies = append(policies, &types.Policy{
+ ID: fmt.Sprintf("policy-res-%d", nr), Name: fmt.Sprintf("Resource Policy %d", nr), Enabled: true,
+ SourcePostureChecks: []string{"posture-check-ver"},
+ Rules: []*types.PolicyRule{{
+ ID: fmt.Sprintf("rule-res-%d", nr), Name: fmt.Sprintf("Allow Resource %d", nr), Enabled: true,
+ Action: types.PolicyTrafficActionAccept, Protocol: types.PolicyRuleProtocolALL, Bidirectional: true,
+ Sources: []string{fmt.Sprintf("group-%d", nr%numGroups)},
+ DestinationResource: types.Resource{ID: resID},
+ }},
+ })
+ }
+
+ account := &types.Account{
+ Id: "test-account",
+ Peers: peers,
+ Groups: groups,
+ Policies: policies,
+ Routes: routes,
+ Users: map[string]*types.User{
+ "user-admin": {Id: "user-admin", Role: types.UserRoleAdmin, IsServiceUser: false, AccountID: "test-account"},
+ },
+ Network: &types.Network{
+ Identifier: "net-test", Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}, Serial: 1,
+ },
+ DNSSettings: types.DNSSettings{DisabledManagementGroups: []string{}},
+ NameServerGroups: map[string]*nbdns.NameServerGroup{
+ "ns-group-main": {
+ ID: "ns-group-main", Name: "Main NS", Enabled: true, Groups: []string{"group-all"},
+ NameServers: []nbdns.NameServer{{IP: netip.MustParseAddr("8.8.8.8"), NSType: nbdns.UDPNameServerType, Port: 53}},
+ },
+ },
+ PostureChecks: []*posture.Checks{
+ {ID: "posture-check-ver", Name: "Check version", Checks: posture.ChecksDefinition{
+ NBVersionCheck: &posture.NBVersionCheck{MinVersion: "0.26.0"},
+ }},
+ },
+ NetworkResources: networkResources,
+ Networks: networksList,
+ NetworkRouters: networkRouters,
+ Settings: &types.Settings{PeerLoginExpirationEnabled: true, PeerLoginExpiration: 1 * time.Hour},
+ }
+
+ for _, p := range account.Policies {
+ p.AccountID = account.Id
+ }
+ for _, r := range account.Routes {
+ r.AccountID = account.Id
+ }
+
+ validatedPeers := make(map[string]struct{}, numPeers)
+ for i := range numPeers {
+ peerID := fmt.Sprintf("peer-%d", i)
+ if i != numPeers-1 {
+ validatedPeers[peerID] = struct{}{}
+ }
+ }
+
+ return account, validatedPeers
+}
+
+// componentsNetworkMap is a convenience wrapper for GetPeerNetworkMapFromComponents.
+func componentsNetworkMap(account *types.Account, peerID string, validatedPeers map[string]struct{}) *types.NetworkMap {
+ return account.GetPeerNetworkMapFromComponents(
+ context.Background(), peerID, nbdns.CustomZone{}, nil,
+ validatedPeers, account.GetResourcePoliciesMap(), account.GetResourceRoutersMap(),
+ nil, account.GetActiveGroupUsers(),
+ )
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 1. PEER VISIBILITY & GROUPS
+// ──────────────────────────────────────────────────────────────────────────────
+
+func TestComponents_PeerVisibility(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(100, 5)
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+ assert.Equal(t, len(validatedPeers)-1-len(nm.OfflinePeers), len(nm.Peers), "peer should see all other validated non-expired peers")
+}
+
+func TestComponents_PeerDoesNotSeeItself(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(50, 5)
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+ for _, p := range nm.Peers {
+ assert.NotEqual(t, "peer-0", p.ID, "peer should not see itself")
+ }
+}
+
+func TestComponents_IntraGroupConnectivity(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(20, 2)
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ peerIDs := make(map[string]bool, len(nm.Peers))
+ for _, p := range nm.Peers {
+ peerIDs[p.ID] = true
+ }
+ assert.True(t, peerIDs["peer-5"], "peer-0 should see peer-5 from same group")
+}
+
+func TestComponents_CrossGroupConnectivity(t *testing.T) {
+ // Without default policy, only per-group policies provide connectivity
+ account, validatedPeers := scalableTestAccountWithoutDefaultPolicy(20, 2)
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ peerIDs := make(map[string]bool, len(nm.Peers))
+ for _, p := range nm.Peers {
+ peerIDs[p.ID] = true
+ }
+ assert.True(t, peerIDs["peer-10"], "peer-0 should see peer-10 from cross-group policy")
+}
+
+func TestComponents_BidirectionalPolicy(t *testing.T) {
+ // Without default policy so bidirectional visibility comes only from per-group policies
+ account, validatedPeers := scalableTestAccountWithoutDefaultPolicy(100, 5)
+ nm0 := componentsNetworkMap(account, "peer-0", validatedPeers)
+ nm20 := componentsNetworkMap(account, "peer-20", validatedPeers)
+ require.NotNil(t, nm0)
+ require.NotNil(t, nm20)
+
+ peer0SeesPeer20 := false
+ for _, p := range nm0.Peers {
+ if p.ID == "peer-20" {
+ peer0SeesPeer20 = true
+ }
+ }
+ peer20SeesPeer0 := false
+ for _, p := range nm20.Peers {
+ if p.ID == "peer-0" {
+ peer20SeesPeer0 = true
+ }
+ }
+ assert.True(t, peer0SeesPeer20, "peer-0 should see peer-20 via bidirectional policy")
+ assert.True(t, peer20SeesPeer0, "peer-20 should see peer-0 via bidirectional policy")
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 2. PEER EXPIRATION & ACCOUNT SETTINGS
+// ──────────────────────────────────────────────────────────────────────────────
+
+func TestComponents_ExpiredPeerInOfflineList(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(100, 5)
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ offlineIDs := make(map[string]bool, len(nm.OfflinePeers))
+ for _, p := range nm.OfflinePeers {
+ offlineIDs[p.ID] = true
+ }
+ assert.True(t, offlineIDs["peer-98"], "expired peer should be in OfflinePeers")
+ for _, p := range nm.Peers {
+ assert.NotEqual(t, "peer-98", p.ID, "expired peer should not be in active Peers")
+ }
+}
+
+func TestComponents_ExpirationDisabledSetting(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(100, 5)
+ account.Settings.PeerLoginExpirationEnabled = false
+
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ peerIDs := make(map[string]bool, len(nm.Peers))
+ for _, p := range nm.Peers {
+ peerIDs[p.ID] = true
+ }
+ assert.True(t, peerIDs["peer-98"], "with expiration disabled, peer-98 should be in active Peers")
+}
+
+func TestComponents_LoginExpiration_PeerLevel(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(20, 2)
+ account.Settings.PeerLoginExpirationEnabled = true
+ account.Settings.PeerLoginExpiration = 1 * time.Hour
+
+ pastLogin := time.Now().Add(-2 * time.Hour)
+ account.Peers["peer-5"].LastLogin = &pastLogin
+ account.Peers["peer-5"].LoginExpirationEnabled = true
+
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ offlineIDs := make(map[string]bool, len(nm.OfflinePeers))
+ for _, p := range nm.OfflinePeers {
+ offlineIDs[p.ID] = true
+ }
+ assert.True(t, offlineIDs["peer-5"], "login-expired peer should be in OfflinePeers")
+ for _, p := range nm.Peers {
+ assert.NotEqual(t, "peer-5", p.ID, "login-expired peer should not be in active Peers")
+ }
+}
+
+func TestComponents_NetworkSerial(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(50, 5)
+ account.Network.Serial = 42
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+ assert.Equal(t, uint64(42), nm.Network.Serial, "network serial should match")
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 3. NON-VALIDATED PEERS
+// ──────────────────────────────────────────────────────────────────────────────
+
+func TestComponents_NonValidatedPeerExcluded(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(100, 5)
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+ for _, p := range nm.Peers {
+ assert.NotEqual(t, "peer-99", p.ID, "non-validated peer should not appear in Peers")
+ }
+ for _, p := range nm.OfflinePeers {
+ assert.NotEqual(t, "peer-99", p.ID, "non-validated peer should not appear in OfflinePeers")
+ }
+}
+
+func TestComponents_NonValidatedTargetPeerGetsEmptyMap(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(100, 5)
+ nm := componentsNetworkMap(account, "peer-99", validatedPeers)
+ require.NotNil(t, nm)
+ assert.Empty(t, nm.Peers)
+ assert.Empty(t, nm.FirewallRules)
+}
+
+func TestComponents_NonExistentPeerGetsEmptyMap(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(100, 5)
+ nm := componentsNetworkMap(account, "peer-does-not-exist", validatedPeers)
+ require.NotNil(t, nm)
+ assert.Empty(t, nm.Peers)
+ assert.Empty(t, nm.FirewallRules)
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 4. POLICIES & FIREWALL RULES
+// ──────────────────────────────────────────────────────────────────────────────
+
+func TestComponents_FirewallRulesGenerated(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(100, 5)
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+ assert.NotEmpty(t, nm.FirewallRules, "should have firewall rules from policies")
+}
+
+func TestComponents_DropPolicyGeneratesDropRules(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(100, 5)
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ hasDropRule := false
+ for _, rule := range nm.FirewallRules {
+ if rule.Action == string(types.PolicyTrafficActionDrop) {
+ hasDropRule = true
+ break
+ }
+ }
+ assert.True(t, hasDropRule, "should have at least one drop firewall rule")
+}
+
+func TestComponents_DisabledPolicyIgnored(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(50, 2)
+ for _, p := range account.Policies {
+ p.Enabled = false
+ }
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+ assert.Empty(t, nm.Peers, "disabled policies should yield no peers")
+ assert.Empty(t, nm.FirewallRules, "disabled policies should yield no firewall rules")
+}
+
+func TestComponents_PortPolicy(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(50, 2)
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ has8080, has5432 := false, false
+ for _, rule := range nm.FirewallRules {
+ if rule.Port == "8080" {
+ has8080 = true
+ }
+ if rule.Port == "5432" {
+ has5432 = true
+ }
+ }
+ assert.True(t, has8080, "should have firewall rule for port 8080")
+ assert.True(t, has5432, "should have firewall rule for port 5432 (drop policy)")
+}
+
+func TestComponents_PortRangePolicy(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(50, 2)
+ account.Peers["peer-0"].Meta.WtVersion = "0.50.0"
+
+ account.Policies = append(account.Policies, &types.Policy{
+ ID: "policy-port-range", Name: "Port Range", Enabled: true, AccountID: "test-account",
+ Rules: []*types.PolicyRule{{
+ ID: "rule-port-range", Name: "Port Range Rule", Enabled: true,
+ Action: types.PolicyTrafficActionAccept, Protocol: types.PolicyRuleProtocolTCP,
+ Bidirectional: true,
+ PortRanges: []types.RulePortRange{{Start: 8000, End: 9000}},
+ Sources: []string{"group-0"}, Destinations: []string{"group-1"},
+ }},
+ })
+
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ hasPortRange := false
+ for _, rule := range nm.FirewallRules {
+ if rule.PortRange.Start == 8000 && rule.PortRange.End == 9000 {
+ hasPortRange = true
+ break
+ }
+ }
+ assert.True(t, hasPortRange, "should have firewall rule with port range 8000-9000")
+}
+
+func TestComponents_FirewallRuleDirection(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(50, 2)
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ hasIn, hasOut := false, false
+ for _, rule := range nm.FirewallRules {
+ if rule.Direction == types.FirewallRuleDirectionIN {
+ hasIn = true
+ }
+ if rule.Direction == types.FirewallRuleDirectionOUT {
+ hasOut = true
+ }
+ }
+ assert.True(t, hasIn, "should have inbound firewall rules")
+ assert.True(t, hasOut, "should have outbound firewall rules")
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 5. ROUTES
+// ──────────────────────────────────────────────────────────────────────────────
+
+func TestComponents_RoutesIncluded(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(100, 5)
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+ assert.NotEmpty(t, nm.Routes, "should have routes")
+}
+
+func TestComponents_DisabledRouteExcluded(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(50, 2)
+ for _, r := range account.Routes {
+ r.Enabled = false
+ }
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+ for _, r := range nm.Routes {
+ assert.True(t, r.Enabled, "only enabled routes should appear")
+ }
+}
+
+func TestComponents_RoutesFirewallRulesForACG(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(100, 5)
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+ assert.NotEmpty(t, nm.RoutesFirewallRules, "should have route firewall rules for access-controlled routes")
+}
+
+func TestComponents_HARouteDeduplication(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(50, 5)
+
+ haNetwork := netip.MustParsePrefix("172.16.0.0/16")
+ account.Routes["route-ha-1"] = &route.Route{
+ ID: "route-ha-1", Network: haNetwork, PeerID: "peer-10",
+ Peer: account.Peers["peer-10"].Key, Enabled: true, Metric: 100,
+ Groups: []string{"group-all"}, PeerGroups: []string{"group-0"}, AccountID: "test-account",
+ }
+ account.Routes["route-ha-2"] = &route.Route{
+ ID: "route-ha-2", Network: haNetwork, PeerID: "peer-20",
+ Peer: account.Peers["peer-20"].Key, Enabled: true, Metric: 200,
+ Groups: []string{"group-all"}, PeerGroups: []string{"group-1"}, AccountID: "test-account",
+ }
+
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ haRoutes := 0
+ for _, r := range nm.Routes {
+ if r.Network == haNetwork {
+ haRoutes++
+ }
+ }
+ // Components deduplicates HA routes with the same HA unique ID, returning one entry per HA group
+ assert.Equal(t, 1, haRoutes, "HA routes with same network should be deduplicated into one entry")
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 6. NETWORK RESOURCES & ROUTERS
+// ──────────────────────────────────────────────────────────────────────────────
+
+func TestComponents_NetworkResourceRoutes_RouterPeer(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(100, 5)
+
+ var routerPeerID string
+ for _, nr := range account.NetworkRouters {
+ routerPeerID = nr.Peer
+ break
+ }
+ require.NotEmpty(t, routerPeerID)
+
+ nm := componentsNetworkMap(account, routerPeerID, validatedPeers)
+ require.NotNil(t, nm)
+ assert.NotEmpty(t, nm.Peers, "router peer should see source peers")
+}
+
+func TestComponents_NetworkResourceRoutes_SourcePeerSeesRouterPeer(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(100, 5)
+
+ var routerPeerID string
+ for _, nr := range account.NetworkRouters {
+ routerPeerID = nr.Peer
+ break
+ }
+ require.NotEmpty(t, routerPeerID)
+
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ peerIDs := make(map[string]bool, len(nm.Peers))
+ for _, p := range nm.Peers {
+ peerIDs[p.ID] = true
+ }
+ assert.True(t, peerIDs[routerPeerID], "source peer should see router peer for network resource")
+}
+
+func TestComponents_DisabledNetworkResourceIgnored(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(50, 5)
+ for _, nr := range account.NetworkResources {
+ nr.Enabled = false
+ }
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+ assert.NotNil(t, nm.Network)
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 7. POSTURE CHECKS
+// ──────────────────────────────────────────────────────────────────────────────
+
+func TestComponents_PostureCheckFiltering_PassingPeer(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(100, 5)
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+ assert.NotEmpty(t, nm.Routes, "passing peer should have routes including resource routes")
+}
+
+func TestComponents_PostureCheckFiltering_FailingPeer(t *testing.T) {
+ // peer-0 has version 0.40.0 (passes posture check >= 0.26.0)
+ // peer-1 has version 0.25.0 (fails posture check >= 0.26.0)
+ // Resource policies require posture-check-ver, so the failing peer
+ // should not see the router peer for those resources.
+ account, validatedPeers := scalableTestAccountWithoutDefaultPolicy(100, 5)
+
+ nm0 := componentsNetworkMap(account, "peer-0", validatedPeers)
+ nm1 := componentsNetworkMap(account, "peer-1", validatedPeers)
+ require.NotNil(t, nm0)
+ require.NotNil(t, nm1)
+
+ // The passing peer should have more peers visible (including resource router peers)
+ // than the failing peer, because the failing peer is excluded from resource policies.
+ assert.Greater(t, len(nm0.Peers), len(nm1.Peers),
+ "passing peer (0.40.0) should see more peers than failing peer (0.25.0) due to posture-gated resource policies")
+}
+
+func TestComponents_MultiplePostureChecks(t *testing.T) {
+ account, validatedPeers := scalableTestAccountWithoutDefaultPolicy(50, 2)
+
+ // Keep only the posture-gated policy — remove per-group policies so connectivity is isolated
+ account.Policies = []*types.Policy{}
+
+ // Set kernel version on peers so the OS posture check can evaluate
+ for _, p := range account.Peers {
+ p.Meta.KernelVersion = "5.15.0"
+ }
+
+ account.PostureChecks = append(account.PostureChecks, &posture.Checks{
+ ID: "posture-check-os", Name: "Check OS",
+ Checks: posture.ChecksDefinition{
+ OSVersionCheck: &posture.OSVersionCheck{Linux: &posture.MinKernelVersionCheck{MinKernelVersion: "0.0.1"}},
+ },
+ })
+ account.Policies = append(account.Policies, &types.Policy{
+ ID: "policy-multi-posture", Name: "Multi Posture", Enabled: true, AccountID: "test-account",
+ SourcePostureChecks: []string{"posture-check-ver", "posture-check-os"},
+ Rules: []*types.PolicyRule{{
+ ID: "rule-multi-posture", Name: "Multi Check Rule", Enabled: true,
+ Action: types.PolicyTrafficActionAccept, Protocol: types.PolicyRuleProtocolALL,
+ Bidirectional: true,
+ Sources: []string{"group-0"}, Destinations: []string{"group-1"},
+ }},
+ })
+
+ // peer-0 (0.40.0, kernel 5.15.0) passes both checks, should see group-1 peers
+ nm0 := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm0)
+ assert.NotEmpty(t, nm0.Peers, "peer passing both posture checks should see destination peers")
+
+ // peer-1 (0.25.0, kernel 5.15.0) fails version check, should NOT see group-1 peers
+ nm1 := componentsNetworkMap(account, "peer-1", validatedPeers)
+ require.NotNil(t, nm1)
+ assert.Empty(t, nm1.Peers,
+ "peer failing posture check should see no peers when posture-gated policy is the only connectivity")
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 8. DNS
+// ──────────────────────────────────────────────────────────────────────────────
+
+func TestComponents_DNSConfigEnabled(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(100, 5)
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+ assert.True(t, nm.DNSConfig.ServiceEnable, "DNS should be enabled")
+ assert.NotEmpty(t, nm.DNSConfig.NameServerGroups, "should have nameserver groups")
+}
+
+func TestComponents_DNSDisabledByManagementGroup(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(100, 5)
+ account.DNSSettings.DisabledManagementGroups = []string{"group-all"}
+
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+ assert.False(t, nm.DNSConfig.ServiceEnable, "DNS should be disabled for peer in disabled group")
+}
+
+func TestComponents_DNSNameServerGroupDistribution(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(20, 2)
+ account.NameServerGroups["ns-group-0"] = &nbdns.NameServerGroup{
+ ID: "ns-group-0", Name: "Group 0 NS", Enabled: true, Groups: []string{"group-0"},
+ NameServers: []nbdns.NameServer{{IP: netip.MustParseAddr("1.1.1.1"), NSType: nbdns.UDPNameServerType, Port: 53}},
+ }
+
+ nm0 := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm0)
+ hasGroup0NS := false
+ for _, ns := range nm0.DNSConfig.NameServerGroups {
+ if ns.ID == "ns-group-0" {
+ hasGroup0NS = true
+ }
+ }
+ assert.True(t, hasGroup0NS, "peer-0 in group-0 should receive ns-group-0")
+
+ nm10 := componentsNetworkMap(account, "peer-10", validatedPeers)
+ require.NotNil(t, nm10)
+ hasGroup0NSForPeer10 := false
+ for _, ns := range nm10.DNSConfig.NameServerGroups {
+ if ns.ID == "ns-group-0" {
+ hasGroup0NSForPeer10 = true
+ }
+ }
+ assert.False(t, hasGroup0NSForPeer10, "peer-10 in group-1 should NOT receive ns-group-0")
+}
+
+func TestComponents_DNSCustomZone(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(20, 2)
+ customZone := nbdns.CustomZone{
+ Domain: "netbird.cloud.",
+ Records: []nbdns.SimpleRecord{
+ {Name: "peer0.netbird.cloud.", Type: 1, Class: "IN", TTL: 300, RData: account.Peers["peer-0"].IP.String()},
+ {Name: "peer1.netbird.cloud.", Type: 1, Class: "IN", TTL: 300, RData: account.Peers["peer-1"].IP.String()},
+ },
+ }
+
+ nm := account.GetPeerNetworkMapFromComponents(
+ context.Background(), "peer-0", customZone, nil,
+ validatedPeers, account.GetResourcePoliciesMap(), account.GetResourceRoutersMap(),
+ nil, account.GetActiveGroupUsers(),
+ )
+ require.NotNil(t, nm)
+ assert.True(t, nm.DNSConfig.ServiceEnable)
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 9. SSH
+// ──────────────────────────────────────────────────────────────────────────────
+
+func TestComponents_SSHPolicy(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(20, 2)
+ account.Groups["ssh-users"] = &types.Group{ID: "ssh-users", Name: "SSH Users", Peers: []string{}}
+ account.Policies = append(account.Policies, &types.Policy{
+ ID: "policy-ssh", Name: "SSH Access", Enabled: true, AccountID: "test-account",
+ Rules: []*types.PolicyRule{{
+ ID: "rule-ssh", Name: "Allow SSH", Enabled: true,
+ Action: types.PolicyTrafficActionAccept, Protocol: types.PolicyRuleProtocolNetbirdSSH,
+ Bidirectional: false,
+ Sources: []string{"group-0"}, Destinations: []string{"group-1"},
+ AuthorizedGroups: map[string][]string{"ssh-users": {"root"}},
+ }},
+ })
+
+ nm := componentsNetworkMap(account, "peer-10", validatedPeers)
+ require.NotNil(t, nm)
+ assert.True(t, nm.EnableSSH, "SSH should be enabled for destination peer of SSH policy")
+}
+
+func TestComponents_SSHNotEnabledWithoutPolicy(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(20, 2)
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+ assert.False(t, nm.EnableSSH, "SSH should not be enabled without SSH policy")
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 10. CROSS-PEER CONSISTENCY
+// ──────────────────────────────────────────────────────────────────────────────
+
+// TestComponents_AllPeersGetValidMaps verifies that every validated peer gets a
+// non-nil map with a consistent network serial and non-empty peer list.
+func TestComponents_AllPeersGetValidMaps(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(50, 5)
+ for peerID := range account.Peers {
+ if _, validated := validatedPeers[peerID]; !validated {
+ continue
+ }
+ nm := componentsNetworkMap(account, peerID, validatedPeers)
+ require.NotNil(t, nm, "network map should not be nil for %s", peerID)
+ assert.Equal(t, account.Network.Serial, nm.Network.Serial, "serial mismatch for %s", peerID)
+ assert.NotEmpty(t, nm.Peers, "validated peer %s should see other peers", peerID)
+ }
+}
+
+// TestComponents_LargeScaleMapGeneration verifies that components can generate maps
+// at larger scales without errors and with consistent output.
+func TestComponents_LargeScaleMapGeneration(t *testing.T) {
+ scales := []struct{ peers, groups int }{
+ {500, 20},
+ {1000, 50},
+ }
+ for _, s := range scales {
+ t.Run(fmt.Sprintf("%dpeers_%dgroups", s.peers, s.groups), func(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(s.peers, s.groups)
+ testPeers := []string{"peer-0", fmt.Sprintf("peer-%d", s.peers/4), fmt.Sprintf("peer-%d", s.peers/2)}
+ for _, peerID := range testPeers {
+ nm := componentsNetworkMap(account, peerID, validatedPeers)
+ require.NotNil(t, nm, "network map should not be nil for %s", peerID)
+ assert.NotEmpty(t, nm.Peers, "peer %s should see other peers at scale", peerID)
+ assert.NotEmpty(t, nm.Routes, "peer %s should have routes at scale", peerID)
+ assert.Equal(t, account.Network.Serial, nm.Network.Serial, "serial mismatch for %s", peerID)
+ }
+ })
+ }
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 11. PEER-AS-RESOURCE POLICIES
+// ──────────────────────────────────────────────────────────────────────────────
+
+// TestComponents_PeerAsSourceResource verifies that a policy with SourceResource.Type=Peer
+// targets only that specific peer as the source.
+func TestComponents_PeerAsSourceResource(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(20, 2)
+
+ account.Policies = append(account.Policies, &types.Policy{
+ ID: "policy-peer-src", Name: "Peer Source Resource", Enabled: true, AccountID: "test-account",
+ Rules: []*types.PolicyRule{{
+ ID: "rule-peer-src", Name: "Peer Source Rule", Enabled: true,
+ Action: types.PolicyTrafficActionAccept,
+ Protocol: types.PolicyRuleProtocolTCP,
+ Bidirectional: true,
+ Ports: []string{"443"},
+ SourceResource: types.Resource{ID: "peer-0", Type: types.ResourceTypePeer},
+ Destinations: []string{"group-1"},
+ }},
+ })
+
+ // peer-0 is the source resource, should see group-1 peers
+ nm0 := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm0)
+
+ has443 := false
+ for _, rule := range nm0.FirewallRules {
+ if rule.Port == "443" {
+ has443 = true
+ break
+ }
+ }
+ assert.True(t, has443, "peer-0 as source resource should have port 443 rule")
+}
+
+// TestComponents_PeerAsDestinationResource verifies that a policy with DestinationResource.Type=Peer
+// targets only that specific peer as the destination.
+func TestComponents_PeerAsDestinationResource(t *testing.T) {
+ account, validatedPeers := scalableTestAccountWithoutDefaultPolicy(20, 2)
+
+ account.Policies = append(account.Policies, &types.Policy{
+ ID: "policy-peer-dst", Name: "Peer Dest Resource", Enabled: true, AccountID: "test-account",
+ Rules: []*types.PolicyRule{{
+ ID: "rule-peer-dst", Name: "Peer Dest Rule", Enabled: true,
+ Action: types.PolicyTrafficActionAccept,
+ Protocol: types.PolicyRuleProtocolTCP,
+ Bidirectional: true,
+ Ports: []string{"443"},
+ Sources: []string{"group-0"},
+ DestinationResource: types.Resource{ID: "peer-15", Type: types.ResourceTypePeer},
+ }},
+ })
+
+ // peer-0 is in group-0 (source), should see peer-15 as destination
+ nm0 := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm0)
+
+ peerIDs := make(map[string]bool, len(nm0.Peers))
+ for _, p := range nm0.Peers {
+ peerIDs[p.ID] = true
+ }
+ assert.True(t, peerIDs["peer-15"], "peer-0 should see peer-15 via peer-as-destination-resource policy")
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 12. MULTIPLE RULES PER POLICY
+// ──────────────────────────────────────────────────────────────────────────────
+
+// TestComponents_MultipleRulesPerPolicy verifies a policy with multiple rules generates
+// firewall rules for each.
+func TestComponents_MultipleRulesPerPolicy(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(20, 2)
+
+ account.Policies = append(account.Policies, &types.Policy{
+ ID: "policy-multi-rule", Name: "Multi Rule Policy", Enabled: true, AccountID: "test-account",
+ Rules: []*types.PolicyRule{
+ {
+ ID: "rule-http", Name: "Allow HTTP", Enabled: true,
+ Action: types.PolicyTrafficActionAccept, Protocol: types.PolicyRuleProtocolTCP,
+ Bidirectional: true, Ports: []string{"80"},
+ Sources: []string{"group-0"}, Destinations: []string{"group-1"},
+ },
+ {
+ ID: "rule-https", Name: "Allow HTTPS", Enabled: true,
+ Action: types.PolicyTrafficActionAccept, Protocol: types.PolicyRuleProtocolTCP,
+ Bidirectional: true, Ports: []string{"443"},
+ Sources: []string{"group-0"}, Destinations: []string{"group-1"},
+ },
+ },
+ })
+
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ has80, has443 := false, false
+ for _, rule := range nm.FirewallRules {
+ if rule.Port == "80" {
+ has80 = true
+ }
+ if rule.Port == "443" {
+ has443 = true
+ }
+ }
+ assert.True(t, has80, "should have firewall rule for port 80 from first rule")
+ assert.True(t, has443, "should have firewall rule for port 443 from second rule")
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 13. SSH AUTHORIZED USERS CONTENT
+// ──────────────────────────────────────────────────────────────────────────────
+
+// TestComponents_SSHAuthorizedUsersContent verifies that SSH policies populate
+// the AuthorizedUsers map with the correct users and machine mappings.
+func TestComponents_SSHAuthorizedUsersContent(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(20, 2)
+
+ account.Users["user-dev"] = &types.User{Id: "user-dev", Role: types.UserRoleUser, AccountID: "test-account", AutoGroups: []string{"ssh-users"}}
+ account.Groups["ssh-users"] = &types.Group{ID: "ssh-users", Name: "SSH Users", Peers: []string{}}
+
+ account.Policies = append(account.Policies, &types.Policy{
+ ID: "policy-ssh", Name: "SSH Access", Enabled: true, AccountID: "test-account",
+ Rules: []*types.PolicyRule{{
+ ID: "rule-ssh", Name: "Allow SSH", Enabled: true,
+ Action: types.PolicyTrafficActionAccept, Protocol: types.PolicyRuleProtocolNetbirdSSH,
+ Bidirectional: false,
+ Sources: []string{"group-0"}, Destinations: []string{"group-1"},
+ AuthorizedGroups: map[string][]string{"ssh-users": {"root", "admin"}},
+ }},
+ })
+
+ // peer-10 is in group-1 (destination)
+ nm := componentsNetworkMap(account, "peer-10", validatedPeers)
+ require.NotNil(t, nm)
+ assert.True(t, nm.EnableSSH, "SSH should be enabled")
+ assert.NotNil(t, nm.AuthorizedUsers, "AuthorizedUsers should not be nil")
+ assert.NotEmpty(t, nm.AuthorizedUsers, "AuthorizedUsers should have entries")
+
+ // Check that "root" machine user mapping exists
+ _, hasRoot := nm.AuthorizedUsers["root"]
+ _, hasAdmin := nm.AuthorizedUsers["admin"]
+ assert.True(t, hasRoot || hasAdmin, "AuthorizedUsers should contain 'root' or 'admin' machine user mapping")
+}
+
+// TestComponents_SSHLegacyImpliedSSH verifies that a non-SSH ALL protocol policy with
+// SSHEnabled peer implies legacy SSH access.
+func TestComponents_SSHLegacyImpliedSSH(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(20, 2)
+
+ // Enable SSH on the destination peer
+ account.Peers["peer-10"].SSHEnabled = true
+
+ // The default "Allow All" policy with Protocol=ALL + SSHEnabled peer should imply SSH
+ nm := componentsNetworkMap(account, "peer-10", validatedPeers)
+ require.NotNil(t, nm)
+ assert.True(t, nm.EnableSSH, "SSH should be implied by ALL protocol policy with SSHEnabled peer")
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 14. ROUTE DEFAULT PERMIT (no AccessControlGroups)
+// ──────────────────────────────────────────────────────────────────────────────
+
+// TestComponents_RouteDefaultPermit verifies that a route without AccessControlGroups
+// generates default permit firewall rules (0.0.0.0/0 source).
+func TestComponents_RouteDefaultPermit(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(20, 2)
+
+ // Add a route without ACGs — this peer is the routing peer
+ routingPeerID := "peer-5"
+ account.Routes["route-no-acg"] = &route.Route{
+ ID: "route-no-acg", Network: netip.MustParsePrefix("192.168.99.0/24"),
+ PeerID: routingPeerID, Peer: account.Peers[routingPeerID].Key,
+ Enabled: true, Groups: []string{"group-all"}, PeerGroups: []string{"group-0"},
+ AccessControlGroups: []string{},
+ AccountID: "test-account",
+ }
+
+ // The routing peer should get default permit route firewall rules
+ nm := componentsNetworkMap(account, routingPeerID, validatedPeers)
+ require.NotNil(t, nm)
+
+ hasDefaultPermit := false
+ for _, rfr := range nm.RoutesFirewallRules {
+ for _, src := range rfr.SourceRanges {
+ if src == "0.0.0.0/0" || src == "::/0" {
+ hasDefaultPermit = true
+ break
+ }
+ }
+ }
+ assert.True(t, hasDefaultPermit, "route without ACG should have default permit rule with 0.0.0.0/0 source")
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 15. MULTIPLE ROUTERS PER NETWORK
+// ──────────────────────────────────────────────────────────────────────────────
+
+// TestComponents_MultipleRoutersPerNetwork verifies that a network resource
+// with multiple routers provides routes through all available routers.
+func TestComponents_MultipleRoutersPerNetwork(t *testing.T) {
+ account, validatedPeers := scalableTestAccountWithoutDefaultPolicy(20, 2)
+
+ netID := "net-multi-router"
+ resID := "res-multi-router"
+ account.Networks = append(account.Networks, &networkTypes.Network{ID: netID, Name: "Multi Router Network", AccountID: "test-account"})
+ account.NetworkResources = append(account.NetworkResources, &resourceTypes.NetworkResource{
+ ID: resID, NetworkID: netID, AccountID: "test-account", Enabled: true,
+ Address: "multi-svc.netbird.cloud",
+ })
+ account.NetworkRouters = append(account.NetworkRouters,
+ &routerTypes.NetworkRouter{ID: "router-a", NetworkID: netID, Peer: "peer-5", Enabled: true, AccountID: "test-account", Metric: 100},
+ &routerTypes.NetworkRouter{ID: "router-b", NetworkID: netID, Peer: "peer-15", Enabled: true, AccountID: "test-account", Metric: 200},
+ )
+ account.Policies = append(account.Policies, &types.Policy{
+ ID: "policy-multi-router-res", Name: "Multi Router Resource", Enabled: true, AccountID: "test-account",
+ Rules: []*types.PolicyRule{{
+ ID: "rule-multi-router-res", Name: "Allow Multi Router", Enabled: true,
+ Action: types.PolicyTrafficActionAccept, Protocol: types.PolicyRuleProtocolALL, Bidirectional: true,
+ Sources: []string{"group-0"}, DestinationResource: types.Resource{ID: resID},
+ }},
+ })
+
+ // peer-0 is in group-0 (source), should see both router peers
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ peerIDs := make(map[string]bool, len(nm.Peers))
+ for _, p := range nm.Peers {
+ peerIDs[p.ID] = true
+ }
+ assert.True(t, peerIDs["peer-5"], "source peer should see router-a (peer-5)")
+ assert.True(t, peerIDs["peer-15"], "source peer should see router-b (peer-15)")
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 16. PEER-AS-NAMESERVER EXCLUSION
+// ──────────────────────────────────────────────────────────────────────────────
+
+// TestComponents_PeerIsNameserverExcludedFromNSGroup verifies that a peer serving
+// as a nameserver does not receive its own NS group in DNS config.
+func TestComponents_PeerIsNameserverExcludedFromNSGroup(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(20, 2)
+
+ // peer-0 has IP 100.64.0.0 — make it a nameserver
+ nsIP := account.Peers["peer-0"].IP
+ account.NameServerGroups["ns-self"] = &nbdns.NameServerGroup{
+ ID: "ns-self", Name: "Self NS", Enabled: true, Groups: []string{"group-all"},
+ NameServers: []nbdns.NameServer{{IP: netip.AddrFrom4([4]byte{nsIP[0], nsIP[1], nsIP[2], nsIP[3]}), NSType: nbdns.UDPNameServerType, Port: 53}},
+ }
+
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ hasSelfNS := false
+ for _, ns := range nm.DNSConfig.NameServerGroups {
+ if ns.ID == "ns-self" {
+ hasSelfNS = true
+ }
+ }
+ assert.False(t, hasSelfNS, "peer serving as nameserver should NOT receive its own NS group")
+
+ // peer-10 is NOT the nameserver, should receive the NS group
+ nm10 := componentsNetworkMap(account, "peer-10", validatedPeers)
+ require.NotNil(t, nm10)
+ hasNSForPeer10 := false
+ for _, ns := range nm10.DNSConfig.NameServerGroups {
+ if ns.ID == "ns-self" {
+ hasNSForPeer10 = true
+ }
+ }
+ assert.True(t, hasNSForPeer10, "non-nameserver peer should receive the NS group")
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 17. DOMAIN NETWORK RESOURCES
+// ──────────────────────────────────────────────────────────────────────────────
+
+// TestComponents_DomainNetworkResource verifies that domain-based network resources
+// produce routes with the correct domain configuration.
+func TestComponents_DomainNetworkResource(t *testing.T) {
+ account, validatedPeers := scalableTestAccountWithoutDefaultPolicy(20, 2)
+
+ netID := "net-domain"
+ resID := "res-domain"
+ account.Networks = append(account.Networks, &networkTypes.Network{ID: netID, Name: "Domain Network", AccountID: "test-account"})
+ account.NetworkResources = append(account.NetworkResources, &resourceTypes.NetworkResource{
+ ID: resID, NetworkID: netID, AccountID: "test-account", Enabled: true,
+ Address: "api.example.com", Type: "domain",
+ })
+ account.NetworkRouters = append(account.NetworkRouters, &routerTypes.NetworkRouter{
+ ID: "router-domain", NetworkID: netID, Peer: "peer-5", Enabled: true, AccountID: "test-account",
+ })
+ account.Policies = append(account.Policies, &types.Policy{
+ ID: "policy-domain-res", Name: "Domain Resource Policy", Enabled: true, AccountID: "test-account",
+ Rules: []*types.PolicyRule{{
+ ID: "rule-domain-res", Name: "Allow Domain", Enabled: true,
+ Action: types.PolicyTrafficActionAccept, Protocol: types.PolicyRuleProtocolALL, Bidirectional: true,
+ Sources: []string{"group-0"}, DestinationResource: types.Resource{ID: resID},
+ }},
+ })
+
+ // peer-0 is source, should get route to the domain resource via peer-5
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ peerIDs := make(map[string]bool, len(nm.Peers))
+ for _, p := range nm.Peers {
+ peerIDs[p.ID] = true
+ }
+ assert.True(t, peerIDs["peer-5"], "source peer should see domain resource router peer")
+}
+
+// ──────────────────────────────────────────────────────────────────────────────
+// 18. DISABLED RULE WITHIN ENABLED POLICY
+// ──────────────────────────────────────────────────────────────────────────────
+
+// TestComponents_DisabledRuleInEnabledPolicy verifies that a disabled rule within
+// an enabled policy does not generate firewall rules.
+func TestComponents_DisabledRuleInEnabledPolicy(t *testing.T) {
+ account, validatedPeers := scalableTestAccount(20, 2)
+
+ account.Policies = append(account.Policies, &types.Policy{
+ ID: "policy-mixed-rules", Name: "Mixed Rules", Enabled: true, AccountID: "test-account",
+ Rules: []*types.PolicyRule{
+ {
+ ID: "rule-enabled", Name: "Enabled Rule", Enabled: true,
+ Action: types.PolicyTrafficActionAccept, Protocol: types.PolicyRuleProtocolTCP,
+ Bidirectional: true, Ports: []string{"3000"},
+ Sources: []string{"group-0"}, Destinations: []string{"group-1"},
+ },
+ {
+ ID: "rule-disabled", Name: "Disabled Rule", Enabled: false,
+ Action: types.PolicyTrafficActionAccept, Protocol: types.PolicyRuleProtocolTCP,
+ Bidirectional: true, Ports: []string{"3001"},
+ Sources: []string{"group-0"}, Destinations: []string{"group-1"},
+ },
+ },
+ })
+
+ nm := componentsNetworkMap(account, "peer-0", validatedPeers)
+ require.NotNil(t, nm)
+
+ has3000, has3001 := false, false
+ for _, rule := range nm.FirewallRules {
+ if rule.Port == "3000" {
+ has3000 = true
+ }
+ if rule.Port == "3001" {
+ has3001 = true
+ }
+ }
+ assert.True(t, has3000, "enabled rule should generate firewall rule for port 3000")
+ assert.False(t, has3001, "disabled rule should NOT generate firewall rule for port 3001")
+}
diff --git a/shared/relay/client/client.go b/shared/relay/client/client.go
index ed1b63435..b10b05617 100644
--- a/shared/relay/client/client.go
+++ b/shared/relay/client/client.go
@@ -333,7 +333,7 @@ func (c *Client) connect(ctx context.Context) (*RelayAddr, error) {
dialers := c.getDialers()
rd := dialer.NewRaceDial(c.log, dialer.DefaultConnectionTimeout, c.connectionURL, dialers...)
- conn, err := rd.Dial()
+ conn, err := rd.Dial(ctx)
if err != nil {
return nil, err
}
diff --git a/shared/relay/client/dialer/race_dialer.go b/shared/relay/client/dialer/race_dialer.go
index 0550fc63e..34359d17e 100644
--- a/shared/relay/client/dialer/race_dialer.go
+++ b/shared/relay/client/dialer/race_dialer.go
@@ -40,10 +40,10 @@ func NewRaceDial(log *log.Entry, connectionTimeout time.Duration, serverURL stri
}
}
-func (r *RaceDial) Dial() (net.Conn, error) {
+func (r *RaceDial) Dial(ctx context.Context) (net.Conn, error) {
connChan := make(chan dialResult, len(r.dialerFns))
winnerConn := make(chan net.Conn, 1)
- abortCtx, abort := context.WithCancel(context.Background())
+ abortCtx, abort := context.WithCancel(ctx)
defer abort()
for _, dfn := range r.dialerFns {
diff --git a/shared/relay/client/dialer/race_dialer_test.go b/shared/relay/client/dialer/race_dialer_test.go
index d216ec5e7..aa18df578 100644
--- a/shared/relay/client/dialer/race_dialer_test.go
+++ b/shared/relay/client/dialer/race_dialer_test.go
@@ -78,7 +78,7 @@ func TestRaceDialEmptyDialers(t *testing.T) {
serverURL := "test.server.com"
rd := NewRaceDial(logger, DefaultConnectionTimeout, serverURL)
- conn, err := rd.Dial()
+ conn, err := rd.Dial(context.Background())
if err == nil {
t.Errorf("Expected an error with empty dialers, got nil")
}
@@ -104,7 +104,7 @@ func TestRaceDialSingleSuccessfulDialer(t *testing.T) {
}
rd := NewRaceDial(logger, DefaultConnectionTimeout, serverURL, mockDialer)
- conn, err := rd.Dial()
+ conn, err := rd.Dial(context.Background())
if err != nil {
t.Errorf("Expected no error, got %v", err)
}
@@ -137,7 +137,7 @@ func TestRaceDialMultipleDialersWithOneSuccess(t *testing.T) {
}
rd := NewRaceDial(logger, DefaultConnectionTimeout, serverURL, mockDialer1, mockDialer2)
- conn, err := rd.Dial()
+ conn, err := rd.Dial(context.Background())
if err != nil {
t.Errorf("Expected no error, got %v", err)
}
@@ -160,7 +160,7 @@ func TestRaceDialTimeout(t *testing.T) {
}
rd := NewRaceDial(logger, 3*time.Second, serverURL, mockDialer)
- conn, err := rd.Dial()
+ conn, err := rd.Dial(context.Background())
if err == nil {
t.Errorf("Expected an error, got nil")
}
@@ -188,7 +188,7 @@ func TestRaceDialAllDialersFail(t *testing.T) {
}
rd := NewRaceDial(logger, DefaultConnectionTimeout, serverURL, mockDialer1, mockDialer2)
- conn, err := rd.Dial()
+ conn, err := rd.Dial(context.Background())
if err == nil {
t.Errorf("Expected an error, got nil")
}
@@ -230,7 +230,7 @@ func TestRaceDialFirstSuccessfulDialerWins(t *testing.T) {
}
rd := NewRaceDial(logger, DefaultConnectionTimeout, serverURL, mockDialer1, mockDialer2)
- conn, err := rd.Dial()
+ conn, err := rd.Dial(context.Background())
if err != nil {
t.Errorf("Expected no error, got %v", err)
}