Merge remote-tracking branch 'origin/main' into fix/pkce-flow-session-extend

# Conflicts:
#	shared/management/proto/management.pb.go
This commit is contained in:
Zoltán Papp
2026-10-07 13:54:40 +02:00
213 changed files with 9278 additions and 8536 deletions
+67 -3
View File
@@ -90,8 +90,9 @@ type StatusRecorder interface {
// fallback T-FinalWarningLead dialog (suppressed when the user dismissed
// the first one for the same deadline). Safe for concurrent use.
type Watcher struct {
lead time.Duration
finalLead time.Duration
lead time.Duration
finalLead time.Duration
deadlineOnly bool
mu sync.Mutex
current time.Time
@@ -102,6 +103,7 @@ type Watcher struct {
dismissedAt time.Time // deadline value the user dismissed via Dismiss(); gates fireFinal
closed bool
recorder StatusRecorder
nowFn func() time.Time
}
// New returns a watcher with the package defaults WarningLead and
@@ -122,9 +124,17 @@ func NewWithLeads(lead, final time.Duration, recorder StatusRecorder) *Watcher {
lead: lead,
finalLead: final,
recorder: recorder,
nowFn: time.Now,
}
}
// NewDeadlineOnly returns a watcher that validates and records deadlines but arms no warning timers.
func NewDeadlineOnly(recorder StatusRecorder) *Watcher {
w := New(recorder)
w.deadlineOnly = true
return w
}
// Update sets the latest deadline. Pass the zero time to clear (e.g. when
// a Sync push from the server omits the field because login expiration
// was disabled).
@@ -181,7 +191,7 @@ func (w *Watcher) Update(deadline time.Time) error {
w.finalFiredAt = time.Time{}
w.dismissedAt = time.Time{}
if deadline.After(now) {
if deadline.After(now) && !w.deadlineOnly {
w.armTimerLocked(deadline)
}
recorder := w.recorder
@@ -303,6 +313,11 @@ func (w *Watcher) fire(armedFor time.Time) {
w.mu.Unlock()
return
}
now := w.nowFn()
if isLate(now, armedFor, max(w.finalLead, 0)) {
w.fireLateLocked(armedFor, now)
return
}
w.firedAt = armedFor
recorder := w.recorder
w.mu.Unlock()
@@ -331,6 +346,14 @@ func (w *Watcher) fireFinal(armedFor time.Time) {
log.Infof("auth session final-warning skipped (dismissed by user)")
return
}
now := w.nowFn()
if isLate(now, armedFor, 0) {
w.finalFiredAt = armedFor
w.mu.Unlock()
log.Infof("auth session final-warning skipped for deadline %s (passed %s ago)",
armedFor.Format(time.RFC3339), now.Round(0).Sub(armedFor).Round(time.Second))
return
}
w.finalFiredAt = armedFor
recorder := w.recorder
w.mu.Unlock()
@@ -341,6 +364,39 @@ func (w *Watcher) fireFinal(armedFor time.Time) {
publishWarning(recorder, armedFor, true)
}
// fireLateLocked handles a T-WarningLead callback that fired inside the
// final-warning window: it sends the final warning in its place while the
// deadline has not passed and the user has not dismissed it, so a resume
// with time left still warns. The caller must hold w.mu; this helper
// releases it.
func (w *Watcher) fireLateLocked(armedFor, now time.Time) {
w.firedAt = armedFor
switch {
case w.dismissedAt.Equal(armedFor):
w.mu.Unlock()
log.Infof("auth session expiry soon warning skipped (dismissed by user)")
return
case w.finalFiredAt.Equal(armedFor):
w.mu.Unlock()
log.Infof("auth session expiry soon warning skipped (final warning already fired)")
return
case isLate(now, armedFor, 0):
w.mu.Unlock()
log.Infof("auth session expiry soon warning skipped for deadline %s (passed %s ago)",
armedFor.Format(time.RFC3339), now.Round(0).Sub(armedFor).Round(time.Second))
return
}
w.finalFiredAt = armedFor
recorder := w.recorder
w.mu.Unlock()
if recorder == nil {
return
}
log.Infof("auth session expiry soon warning fired inside the final-warning window, sending final warning for deadline %s",
armedFor.Format(time.RFC3339))
publishWarning(recorder, armedFor, true)
}
// armOneShotLocked schedules cb at fireAt. When fireAt is already in the
// past it dispatches on the next scheduler tick so a state-change recorder
// notification (invoked after w.mu is released) lands first. Caller must
@@ -380,3 +436,11 @@ func publishWarning(recorder StatusRecorder, deadline time.Time, final bool) {
meta,
)
}
// isLate reports whether the wall clock now has already reached armedFor
// minus cutoffLead. The timers run on the monotonic clock, which can stall
// while the host sleeps, so a timer can fire long after the window it was
// armed for.
func isLate(now, armedFor time.Time, cutoffLead time.Duration) bool {
return !now.Round(0).Before(armedFor.Add(-cutoffLead).Round(0))
}
@@ -527,3 +527,201 @@ func TestDismissBeforeUpdateIsNoop(t *testing.T) {
}
t.Fatalf("final-warning did not publish after no-op pre-Update Dismiss, events=%+v", r.snapshot())
}
func TestIsLate(t *testing.T) {
armedFor := time.Date(2026, 10, 1, 12, 0, 0, 0, time.UTC)
lead := 2 * time.Minute
tests := []struct {
name string
now time.Time
cutoffLead time.Duration
want bool
}{
{"before cutoff", armedFor.Add(-3 * time.Minute), lead, false},
{"at cutoff", armedFor.Add(-lead), lead, true},
{"after cutoff", armedFor.Add(-time.Minute), lead, true},
{"zero lead before deadline", armedFor.Add(-time.Second), 0, false},
{"zero lead at deadline", armedFor, 0, true},
{"zero lead after deadline", armedFor.Add(time.Second), 0, true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := isLate(tt.now, armedFor, tt.cutoffLead); got != tt.want {
t.Fatalf("isLate(%s, %s, %s) = %v, want %v", tt.now, armedFor, tt.cutoffLead, got, tt.want)
}
})
}
}
func TestIsLateIgnoresMonotonicReading(t *testing.T) {
now := time.Now()
wallOnly := now.Round(0)
if isLate(now, wallOnly.Add(time.Second), 0) {
t.Fatalf("now with monotonic reading must compare as wall clock before a later wall-only deadline")
}
if !isLate(now, wallOnly, 0) {
t.Fatalf("now with monotonic reading must compare as wall clock at an equal wall-only deadline")
}
}
func TestLateTimerFiring(t *testing.T) {
tests := []struct {
name string
final bool
beforeDl time.Duration
wantWarns int
wantFinals int
}{
{"warning on resume inside window", false, 3 * time.Minute, 1, 0},
{"warning promoted to final inside final window", false, time.Minute, 0, 1},
{"warning skipped past deadline", false, -time.Minute, 0, 0},
{"final on resume before deadline", true, time.Minute, 0, 1},
{"final skipped past deadline", true, -time.Minute, 0, 0},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
r := &fakeRecorder{}
w := New(r)
defer w.Close()
// The deadline is an hour out so the real timers never fire
// during the test; the late callback is invoked directly with an
// injected clock that simulates a resume near the deadline.
d := time.Now().Add(time.Hour).Round(0)
w.nowFn = func() time.Time { return d.Add(-tt.beforeDl) }
if err := w.Update(d); err != nil {
t.Fatalf("Update: %v", err)
}
if tt.final {
w.fireFinal(d)
} else {
w.fire(d)
}
events := r.snapshot()
if got := countWhere(events, event.isWarning); got != tt.wantWarns {
t.Fatalf("expected %d warning publishes, got %d: %+v", tt.wantWarns, got, events)
}
if got := countWhere(events, event.isFinalWarning); got != tt.wantFinals {
t.Fatalf("expected %d final-warning publishes, got %d: %+v", tt.wantFinals, got, events)
}
})
}
}
func TestPromotedFinalWarningIsNotRepeated(t *testing.T) {
r := &fakeRecorder{}
w := New(r)
defer w.Close()
d := time.Now().Add(time.Hour).Round(0)
now := d.Add(-time.Minute)
w.nowFn = func() time.Time { return now }
if err := w.Update(d); err != nil {
t.Fatalf("Update: %v", err)
}
w.fire(d)
// The final timer was suspended too, so it fires even later than the
// warning timer, here still just before the deadline.
now = d.Add(-30 * time.Second)
w.fireFinal(d)
events := r.snapshot()
if got := countWhere(events, event.isFinalWarning); got != 1 {
t.Fatalf("expected exactly 1 final-warning publish, got %d: %+v", got, events)
}
if got := countWhere(events, event.isWarning); got != 0 {
t.Fatalf("expected no regular warning publish, got %d: %+v", got, events)
}
}
func TestPromotionRespectsDismiss(t *testing.T) {
r := &fakeRecorder{}
w := New(r)
defer w.Close()
d := time.Now().Add(time.Hour).Round(0)
w.nowFn = func() time.Time { return d.Add(-time.Minute) }
if err := w.Update(d); err != nil {
t.Fatalf("Update: %v", err)
}
w.Dismiss()
w.fire(d)
events := r.snapshot()
if got := countWhere(events, func(e event) bool { return e.kind == publish }); got != 0 {
t.Fatalf("expected no publish after dismiss, got %d: %+v", got, events)
}
}
func TestPromotionSkippedWhenFinalAlreadyFired(t *testing.T) {
r := &fakeRecorder{}
w := New(r)
defer w.Close()
// Both timers fall in the past after a long suspend and are dispatched
// with a zero delay, so the final callback can run before the warning one.
d := time.Now().Add(time.Hour).Round(0)
w.nowFn = func() time.Time { return d.Add(-time.Minute) }
if err := w.Update(d); err != nil {
t.Fatalf("Update: %v", err)
}
w.fireFinal(d)
w.fire(d)
events := r.snapshot()
if got := countWhere(events, event.isFinalWarning); got != 1 {
t.Fatalf("expected exactly 1 final-warning publish, got %d: %+v", got, events)
}
if got := countWhere(events, event.isWarning); got != 0 {
t.Fatalf("expected no regular warning publish, got %d: %+v", got, events)
}
}
func TestDeadlineOnlyRecordsDeadlineWithoutWarnings(t *testing.T) {
r := &fakeRecorder{}
w := NewDeadlineOnly(r)
defer w.Close()
// With the default leads this deadline would otherwise fire both
// timers on the next tick.
d := time.Now().Add(50 * time.Millisecond).Round(0)
if err := w.Update(d); err != nil {
t.Fatalf("Update: %v", err)
}
if got := r.deadline(); !got.Equal(d) {
t.Fatalf("expected recorder deadline %v, got %v", d, got)
}
time.Sleep(100 * time.Millisecond)
events := r.snapshot()
if got := countWhere(events, func(e event) bool { return e.kind == publish }); got != 0 {
t.Fatalf("expected no publish in deadline-only mode, got %d: %+v", got, events)
}
if w.timer != nil || w.finalTimer != nil {
t.Fatal("expected no timers armed in deadline-only mode")
}
}
func TestDeadlineOnlyStillRejectsOutOfRangeDeadlines(t *testing.T) {
r := &fakeRecorder{}
w := NewDeadlineOnly(r)
defer w.Close()
if err := w.Update(time.Now().Add(time.Hour)); err != nil {
t.Fatalf("Update: %v", err)
}
err := w.Update(time.Now().Add(-maxPastHorizon - time.Hour))
if !errors.Is(err, ErrDeadlineInPast) {
t.Fatalf("expected ErrDeadlineInPast, got %v", err)
}
if got := r.deadline(); !got.IsZero() {
t.Fatalf("expected recorder cleared after rejection, got %v", got)
}
}
+1
View File
@@ -846,6 +846,7 @@ func TestAddConfig_AllFieldsCovered(t *testing.T) {
"ClientCertKeyPair": "non-config: parsed cert pair, not serialized",
"Name": "non-config: profile name is not needed for debug purposes",
"policy": "non-config: in-memory MDM policy snapshot, surfaced via Config.Policy() / GetConfigResponse.MDMManagedFields",
"probing": "non-config: marks a throwaway copy built to be diffed against; never set on a config anyone runs with",
"DebugBundleUploadURL": "sensitive: MDM-provided upload URL may carry credentials or query tokens; kept out of the shared bundle",
}
+4 -88
View File
@@ -42,7 +42,6 @@ import (
dnsconfig "github.com/netbirdio/netbird/client/internal/dns/config"
"github.com/netbirdio/netbird/client/internal/dnsfwd"
"github.com/netbirdio/netbird/client/internal/expose"
"github.com/netbirdio/netbird/client/internal/ingressgw"
"github.com/netbirdio/netbird/client/internal/lazyconn"
"github.com/netbirdio/netbird/client/internal/metrics"
"github.com/netbirdio/netbird/client/internal/netflow"
@@ -262,11 +261,10 @@ type Engine struct {
statusRecorder *peer.Status
firewall firewallManager.Manager
routeManager routemanager.Manager
acl acl.Manager
dnsForwardMgr *dnsfwd.Manager
ingressGatewayMgr *ingressgw.Manager
firewall firewallManager.Manager
routeManager routemanager.Manager
acl acl.Manager
dnsForwardMgr *dnsfwd.Manager
dnsServer dns.Server
@@ -448,13 +446,6 @@ func (e *Engine) stopLocked() {
e.cleanupSSHConfig()
if e.ingressGatewayMgr != nil {
if err := e.ingressGatewayMgr.Close(); err != nil {
log.Warnf("failed to cleanup forward rules: %v", err)
}
e.ingressGatewayMgr = nil
}
if e.srWatcher != nil {
e.srWatcher.Close()
}
@@ -1627,13 +1618,6 @@ func (e *Engine) updateNetworkMap(networkMap *mgmProto.NetworkMap) error {
e.updateDNSForwarder(dnsRouteFeatureFlag, fwdEntries)
done()
// Ingress forward rules
done = e.phase("forward_rules")
if _, err := e.updateForwardRules(networkMap.GetForwardingRules()); err != nil {
log.Errorf("failed to update forward rules, err: %v", err)
}
done()
log.Debugf("got peers update from Management Service, total peers to connect to = %d", len(networkMap.GetRemotePeers()))
done = e.phase("offline_peers")
@@ -2733,74 +2717,6 @@ func (e *Engine) setForwarderCapture(pc device.PacketCapture) {
}
}
func (e *Engine) updateForwardRules(rules []*mgmProto.ForwardingRule) ([]firewallManager.ForwardRule, error) {
if e.firewall == nil {
log.Warn("firewall is disabled, not updating forwarding rules")
return nil, nil
}
if len(rules) == 0 {
if e.ingressGatewayMgr == nil {
return nil, nil
}
err := e.ingressGatewayMgr.Close()
e.ingressGatewayMgr = nil
e.statusRecorder.SetIngressGwMgr(nil)
return nil, err
}
if e.ingressGatewayMgr == nil {
mgr := ingressgw.NewManager(e.firewall)
e.ingressGatewayMgr = mgr
e.statusRecorder.SetIngressGwMgr(mgr)
}
var merr *multierror.Error
forwardingRules := make([]firewallManager.ForwardRule, 0, len(rules))
for _, rule := range rules {
proto, err := acl.ConvertToFirewallProtocol(rule.GetProtocol())
if err != nil {
merr = multierror.Append(merr, fmt.Errorf("failed to convert protocol '%s': %w", rule.GetProtocol(), err))
continue
}
dstPortInfo, err := convertPortInfo(rule.GetDestinationPort())
if err != nil {
merr = multierror.Append(merr, fmt.Errorf("invalid destination port '%v': %w", rule.GetDestinationPort(), err))
continue
}
translateIP, err := convertToIP(rule.GetTranslatedAddress())
if err != nil {
merr = multierror.Append(merr, fmt.Errorf("failed to convert translated address '%s': %w", rule.GetTranslatedAddress(), err))
continue
}
translatePort, err := convertPortInfo(rule.GetTranslatedPort())
if err != nil {
merr = multierror.Append(merr, fmt.Errorf("invalid translate port '%v': %w", rule.GetTranslatedPort(), err))
continue
}
forwardRule := firewallManager.ForwardRule{
Protocol: proto,
DestinationPort: *dstPortInfo,
TranslatedAddress: translateIP,
TranslatedPort: *translatePort,
}
forwardingRules = append(forwardingRules, forwardRule)
}
log.Infof("updating forwarding rules: %d", len(forwardingRules))
if err := e.ingressGatewayMgr.Update(forwardingRules); err != nil {
log.Errorf("failed to update forwarding rules: %v", err)
}
return forwardingRules, nberrors.FormatErrorOrNil(merr)
}
// toExcludedLazyPeers returns the peers that must have an always-active
// connection: those that are not lazy by policy (the per-peer lazy state or the
// account flag, subject to the local override).
+2 -3
View File
@@ -42,7 +42,6 @@ import (
nbcache "github.com/netbirdio/netbird/management/server/cache"
"github.com/netbirdio/netbird/management/server/groups"
"github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator"
"github.com/netbirdio/netbird/management/server/integrations/port_forwarding"
"github.com/netbirdio/netbird/management/server/job"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/settings"
@@ -523,8 +522,8 @@ func startManagement(t *testing.T, dataDir, testFile string) (*grpc.Server, stri
updateManager := update_channel.NewPeersUpdateManager(metrics)
requestBuffer := server.NewAccountRequestBuffer(context.Background(), store)
networkMapController := controller.NewController(context.Background(), store, metrics, updateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(store, peersManager), config, nil)
accountManager, err := server.BuildManager(context.Background(), config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, ia, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
networkMapController := controller.NewController(context.Background(), store, metrics, updateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", manager.NewEphemeralManager(store, peersManager), config, nil)
accountManager, err := server.BuildManager(context.Background(), config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, ia, metrics, settingsMockManager, permissionsManager, false, cacheStore)
if err != nil {
return nil, "", err
}
+7 -5
View File
@@ -1,4 +1,4 @@
//go:build !js
//go:build !js && !android
package internal
@@ -7,10 +7,12 @@ import (
"github.com/netbirdio/netbird/client/internal/peer"
)
// newSessionWatcher returns the real SSO session expiry watcher for every
// non-wasm build. The js/wasm build gets a no-op stub from
// engine_sessionwatch_js.go so the sessionwatch package (and its timer
// machinery) never links into the wasm binary.
// newSessionWatcher returns the real SSO session expiry watcher. The js/wasm
// build gets a no-op stub from engine_sessionwatch_js.go so the sessionwatch
// package (and its timer machinery) never links into the wasm binary; the
// android build gets a deadline-only watcher from
// engine_sessionwatch_android.go because the app schedules the warnings
// itself.
func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher {
return sessionwatch.New(recorder)
}
@@ -0,0 +1,12 @@
//go:build android
package internal
import (
"github.com/netbirdio/netbird/client/internal/auth/sessionwatch"
"github.com/netbirdio/netbird/client/internal/peer"
)
func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher {
return sessionwatch.NewDeadlineOnly(recorder)
}
-111
View File
@@ -1,111 +0,0 @@
package ingressgw
import (
"fmt"
"sync"
"github.com/hashicorp/go-multierror"
log "github.com/sirupsen/logrus"
nberrors "github.com/netbirdio/netbird/client/errors"
firewall "github.com/netbirdio/netbird/client/firewall/manager"
)
type DNATFirewall interface {
AddDNATRule(fwdRule firewall.ForwardRule) (firewall.Rule, error)
DeleteDNATRule(rule firewall.Rule) error
}
type RulePair struct {
firewall.ForwardRule
firewall.Rule
}
type Manager struct {
dnatFirewall DNATFirewall
rules map[firewall.RuleID]RulePair
rulesMu sync.Mutex
}
func NewManager(dnatFirewall DNATFirewall) *Manager {
return &Manager{
dnatFirewall: dnatFirewall,
rules: make(map[firewall.RuleID]RulePair),
}
}
func (h *Manager) Update(forwardRules []firewall.ForwardRule) error {
h.rulesMu.Lock()
defer h.rulesMu.Unlock()
var mErr *multierror.Error
toDelete := make(map[firewall.RuleID]RulePair, len(h.rules))
for id, r := range h.rules {
toDelete[id] = r
}
// Process new/updated rules
for _, fwdRule := range forwardRules {
id := fwdRule.ID()
if _, ok := h.rules[id]; ok {
delete(toDelete, id)
continue
}
rule, err := h.dnatFirewall.AddDNATRule(fwdRule)
if err != nil {
mErr = multierror.Append(mErr, fmt.Errorf("add forward rule '%s': %v", fwdRule.String(), err))
continue
}
if rule == nil {
mErr = multierror.Append(mErr, fmt.Errorf("add forward rule '%s': backend returned no rule", fwdRule.String()))
continue
}
log.Infof("forward rule has been added '%s'", fwdRule)
h.rules[id] = RulePair{
ForwardRule: fwdRule,
Rule: rule,
}
}
// Remove deleted rules
for id, rulePair := range toDelete {
if err := h.dnatFirewall.DeleteDNATRule(rulePair.Rule); err != nil {
mErr = multierror.Append(mErr, fmt.Errorf("failed to delete forward rule '%s': %v", rulePair.ForwardRule.String(), err))
}
log.Infof("forward rule has been deleted '%s'", rulePair.ForwardRule)
delete(h.rules, id)
}
return nberrors.FormatErrorOrNil(mErr)
}
func (h *Manager) Close() error {
h.rulesMu.Lock()
defer h.rulesMu.Unlock()
log.Infof("clean up all (%d) forward rules", len(h.rules))
var mErr *multierror.Error
for _, rule := range h.rules {
if err := h.dnatFirewall.DeleteDNATRule(rule.Rule); err != nil {
mErr = multierror.Append(mErr, fmt.Errorf("failed to delete forward rule '%s': %v", rule, err))
}
}
h.rules = make(map[firewall.RuleID]RulePair)
return nberrors.FormatErrorOrNil(mErr)
}
func (h *Manager) Rules() []firewall.ForwardRule {
h.rulesMu.Lock()
defer h.rulesMu.Unlock()
rules := make([]firewall.ForwardRule, 0, len(h.rules))
for _, rulePair := range h.rules {
rules = append(rules, rulePair.ForwardRule)
}
return rules
}
-281
View File
@@ -1,281 +0,0 @@
package ingressgw
import (
"fmt"
"net/netip"
"testing"
firewall "github.com/netbirdio/netbird/client/firewall/manager"
)
var (
_ firewall.Rule = (*MocFwRule)(nil)
_ DNATFirewall = &MockDNATFirewall{}
)
type MocFwRule struct {
id firewall.RuleID
}
func (m *MocFwRule) ID() firewall.RuleID {
return m.id
}
type MockDNATFirewall struct {
throwError bool
}
func (m *MockDNATFirewall) AddDNATRule(fwdRule firewall.ForwardRule) (firewall.Rule, error) {
if m.throwError {
return nil, fmt.Errorf("moc error")
}
fwRule := &MocFwRule{
id: fwdRule.ID(),
}
return fwRule, nil
}
func (m *MockDNATFirewall) DeleteDNATRule(rule firewall.Rule) error {
if m.throwError {
return fmt.Errorf("moc error")
}
return nil
}
func (m *MockDNATFirewall) forceToThrowErrors() {
m.throwError = true
}
func TestManager_AddRule(t *testing.T) {
fw := &MockDNATFirewall{}
mgr := NewManager(fw)
port, _ := firewall.NewPort(8080)
updates := []firewall.ForwardRule{
{
Protocol: firewall.ProtocolTCP,
DestinationPort: *port,
TranslatedAddress: netip.MustParseAddr("172.16.254.1"),
TranslatedPort: *port,
},
{
Protocol: firewall.ProtocolUDP,
DestinationPort: *port,
TranslatedAddress: netip.MustParseAddr("172.16.254.1"),
TranslatedPort: *port,
}}
if err := mgr.Update(updates); err != nil {
t.Errorf("unexpected error: %v", err)
}
rules := mgr.Rules()
if len(rules) != len(updates) {
t.Errorf("unexpected rules count: %d", len(rules))
}
}
func TestManager_UpdateRule(t *testing.T) {
fw := &MockDNATFirewall{}
mgr := NewManager(fw)
port, _ := firewall.NewPort(8080)
ruleTCP := firewall.ForwardRule{
Protocol: firewall.ProtocolTCP,
DestinationPort: *port,
TranslatedAddress: netip.MustParseAddr("172.16.254.1"),
TranslatedPort: *port,
}
if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err != nil {
t.Errorf("unexpected error: %v", err)
}
ruleUDP := firewall.ForwardRule{
Protocol: firewall.ProtocolUDP,
DestinationPort: *port,
TranslatedAddress: netip.MustParseAddr("172.16.254.2"),
TranslatedPort: *port,
}
if err := mgr.Update([]firewall.ForwardRule{ruleUDP}); err != nil {
t.Errorf("unexpected error: %v", err)
}
rules := mgr.Rules()
if len(rules) != 1 {
t.Errorf("unexpected rules count: %d", len(rules))
}
if rules[0].TranslatedAddress.String() != ruleUDP.TranslatedAddress.String() {
t.Errorf("unexpected rule: %v", rules[0])
}
if rules[0].TranslatedPort.String() != ruleUDP.TranslatedPort.String() {
t.Errorf("unexpected rule: %v", rules[0])
}
if rules[0].DestinationPort.String() != ruleUDP.DestinationPort.String() {
t.Errorf("unexpected rule: %v", rules[0])
}
if rules[0].Protocol != ruleUDP.Protocol {
t.Errorf("unexpected rule: %v", rules[0])
}
}
func TestManager_ExtendRules(t *testing.T) {
fw := &MockDNATFirewall{}
mgr := NewManager(fw)
port, _ := firewall.NewPort(8080)
ruleTCP := firewall.ForwardRule{
Protocol: firewall.ProtocolTCP,
DestinationPort: *port,
TranslatedAddress: netip.MustParseAddr("172.16.254.1"),
TranslatedPort: *port,
}
ruleUDP := firewall.ForwardRule{
Protocol: firewall.ProtocolUDP,
DestinationPort: *port,
TranslatedAddress: netip.MustParseAddr("172.16.254.2"),
TranslatedPort: *port,
}
if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err != nil {
t.Errorf("unexpected error: %v", err)
}
if err := mgr.Update([]firewall.ForwardRule{ruleTCP, ruleUDP}); err != nil {
t.Errorf("unexpected error: %v", err)
}
rules := mgr.Rules()
if len(rules) != 2 {
t.Errorf("unexpected rules count: %d", len(rules))
}
}
func TestManager_UnderlingError(t *testing.T) {
fw := &MockDNATFirewall{}
mgr := NewManager(fw)
port, _ := firewall.NewPort(8080)
ruleTCP := firewall.ForwardRule{
Protocol: firewall.ProtocolTCP,
DestinationPort: *port,
TranslatedAddress: netip.MustParseAddr("172.16.254.1"),
TranslatedPort: *port,
}
ruleUDP := firewall.ForwardRule{
Protocol: firewall.ProtocolUDP,
DestinationPort: *port,
TranslatedAddress: netip.MustParseAddr("172.16.254.2"),
TranslatedPort: *port,
}
if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err != nil {
t.Errorf("unexpected error: %v", err)
}
fw.forceToThrowErrors()
if err := mgr.Update([]firewall.ForwardRule{ruleTCP, ruleUDP}); err == nil {
t.Errorf("expected error")
}
rules := mgr.Rules()
if len(rules) != 1 {
t.Errorf("unexpected rules count: %d", len(rules))
}
}
func TestManager_Cleanup(t *testing.T) {
fw := &MockDNATFirewall{}
mgr := NewManager(fw)
port, _ := firewall.NewPort(8080)
ruleTCP := firewall.ForwardRule{
Protocol: firewall.ProtocolTCP,
DestinationPort: *port,
TranslatedAddress: netip.MustParseAddr("172.16.254.1"),
TranslatedPort: *port,
}
if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err != nil {
t.Errorf("unexpected error: %v", err)
}
if err := mgr.Update([]firewall.ForwardRule{}); err != nil {
t.Errorf("unexpected error: %v", err)
}
rules := mgr.Rules()
if len(rules) != 0 {
t.Errorf("unexpected rules count: %d", len(rules))
}
}
func TestManager_DeleteBrokenRule(t *testing.T) {
fw := &MockDNATFirewall{}
// force to throw errors when Add DNAT Rule
fw.forceToThrowErrors()
mgr := NewManager(fw)
port, _ := firewall.NewPort(8080)
ruleTCP := firewall.ForwardRule{
Protocol: firewall.ProtocolTCP,
DestinationPort: *port,
TranslatedAddress: netip.MustParseAddr("172.16.254.1"),
TranslatedPort: *port,
}
if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err == nil {
t.Errorf("unexpected error: %v", err)
}
rules := mgr.Rules()
if len(rules) != 0 {
t.Errorf("unexpected rules count: %d", len(rules))
}
// simulate that to remove a broken rule
if err := mgr.Update([]firewall.ForwardRule{}); err != nil {
t.Errorf("unexpected error: %v", err)
}
if err := mgr.Close(); err != nil {
t.Errorf("unexpected error: %v", err)
}
}
func TestManager_Close(t *testing.T) {
fw := &MockDNATFirewall{}
mgr := NewManager(fw)
port, _ := firewall.NewPort(8080)
ruleTCP := firewall.ForwardRule{
Protocol: firewall.ProtocolTCP,
DestinationPort: *port,
TranslatedAddress: netip.MustParseAddr("172.16.254.1"),
TranslatedPort: *port,
}
if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err != nil {
t.Errorf("unexpected error: %v", err)
}
if err := mgr.Close(); err != nil {
t.Errorf("unexpected error: %v", err)
}
rules := mgr.Rules()
if len(rules) != 0 {
t.Errorf("unexpected rules count: %d", len(rules))
}
}
-43
View File
@@ -1,43 +0,0 @@
package internal
import (
"errors"
"fmt"
"net"
"net/netip"
firewallManager "github.com/netbirdio/netbird/client/firewall/manager"
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
)
func convertPortInfo(portInfo *mgmProto.PortInfo) (*firewallManager.Port, error) {
if portInfo == nil {
return nil, errors.New("portInfo cannot be nil")
}
if portInfo.GetPort() != 0 {
return firewallManager.NewPort(int(portInfo.GetPort()))
}
if portInfo.GetRange() != nil {
return firewallManager.NewPort(int(portInfo.GetRange().Start), int(portInfo.GetRange().End))
}
return nil, fmt.Errorf("invalid portInfo: %v", portInfo)
}
func convertToIP(rawIP []byte) (netip.Addr, error) {
if rawIP == nil {
return netip.Addr{}, errors.New("input bytes cannot be nil")
}
if len(rawIP) != net.IPv4len && len(rawIP) != net.IPv6len {
return netip.Addr{}, fmt.Errorf("invalid IP length: %d", len(rawIP))
}
if len(rawIP) == net.IPv4len {
return netip.AddrFrom4([4]byte(rawIP)), nil
}
return netip.AddrFrom16([16]byte(rawIP)), nil
}
+16 -12
View File
@@ -828,7 +828,8 @@ func (conn *Conn) evalStatus() ConnStatus {
//
// The result is a tri-state:
// - ConnStatusConnected: all available transports are up
// - ConnStatusPartiallyConnected: relay is up but ICE is still pending/reconnecting
// - ConnStatusPartiallyConnected: one transport carries the traffic and the other does
// not: relay up with ICE down, or ICE up with the shared relay transport down
// - ConnStatusDisconnected: no working transport
func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) {
defer func() {
@@ -845,13 +846,14 @@ func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) {
}
return evalConnStatus(connStatusInputs{
forceRelay: IsForceRelayed(),
peerUsesRelay: conn.workerRelay.IsRelayConnectionSupportedWithPeer(),
relayConnected: conn.statusRelay.Get() == worker.StatusConnected,
remoteSupportsICE: conn.handshaker.RemoteICESupported(),
iceWorkerCreated: iceWorkerCreated,
iceStatusConnecting: conn.statusICE.Get() != worker.StatusDisconnected,
iceInProgress: iceInProgress,
forceRelay: IsForceRelayed(),
peerUsesRelay: conn.workerRelay.IsRelayConnectionSupportedWithPeer(),
relayConnected: conn.statusRelay.Get() == worker.StatusConnected,
relayTransportConnected: conn.workerRelay.IsTransportConnected(),
remoteSupportsICE: conn.handshaker.RemoteICESupported(),
iceWorkerCreated: iceWorkerCreated,
iceStatusConnected: conn.statusICE.Get() == worker.StatusConnected,
iceInProgress: iceInProgress,
})
}
@@ -1060,19 +1062,21 @@ func evalConnStatus(in connStatusInputs) guard.ConnStatus {
return boolToConnStatus(relayUsedAndUp)
}
// ICE counts as "up" when the status is anything other than Disconnected, OR
// when a negotiation is currently in progress (so we don't spam offers while one is in flight).
iceUp := in.iceStatusConnecting || in.iceInProgress
// ICE counts as "running" when either connected or attempting to connect.
iceRunning := in.iceStatusConnected || in.iceInProgress
// Relay side is acceptable if the peer doesn't rely on relay, or relay is connected.
relayOK := !in.peerUsesRelay || in.relayConnected
switch {
case iceUp && relayOK:
case iceRunning && relayOK:
return guard.ConnStatusConnected
case relayUsedAndUp:
// Relay is up but ICE is down — partially connected.
return guard.ConnStatusPartiallyConnected
case in.iceStatusConnected && !in.relayTransportConnected:
// ICE is up and the shared relay transport is down — offers cannot restore it.
return guard.ConnStatusPartiallyConnected
default:
return guard.ConnStatusDisconnected
}
+8 -7
View File
@@ -17,13 +17,14 @@ const (
// tri-state connection classification. Extracted so the decision logic can be unit-tested
// without constructing full Worker/Handshaker objects.
type connStatusInputs struct {
forceRelay bool // NB_FORCE_RELAY or JS/WASM
peerUsesRelay bool // remote peer advertises relay support AND local has relay
relayConnected bool // statusRelay reports Connected (independent of whether peer uses relay)
remoteSupportsICE bool // remote peer sent ICE credentials
iceWorkerCreated bool // local WorkerICE exists (false in force-relay mode)
iceStatusConnecting bool // statusICE is anything other than Disconnected
iceInProgress bool // a negotiation is currently in flight
forceRelay bool // NB_FORCE_RELAY or JS/WASM
peerUsesRelay bool // remote peer advertises relay support AND local has relay
relayConnected bool // statusRelay reports Connected (independent of whether peer uses relay)
relayTransportConnected bool // the relay transport shared by all peers on that server is up
remoteSupportsICE bool // remote peer sent ICE credentials
iceWorkerCreated bool // local WorkerICE exists (false in force-relay mode)
iceStatusConnected bool // statusICE reports Connected
iceInProgress bool // a negotiation is currently in flight
}
// ConnStatus describe the status of a peer's connection
+71 -12
View File
@@ -30,6 +30,21 @@ func TestEvalConnStatus_ForceRelay(t *testing.T) {
},
want: guard.ConnStatusDisconnected,
},
{
name: "force relay, relay up but the shared transport reports down",
in: connStatusInputs{
forceRelay: true,
peerUsesRelay: true,
relayConnected: true,
relayTransportConnected: false,
// The ICE inputs are set so that the force-relay return is the only branch
// that can produce Connected here: without it the peer would fall through to
// relayUsedAndUp and report PartiallyConnected.
remoteSupportsICE: true,
iceWorkerCreated: true,
},
want: guard.ConnStatusConnected,
},
{
name: "force relay, peer does NOT use relay - disconnected forever",
in: connStatusInputs{
@@ -123,24 +138,28 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true
in.relayConnected = true
in.iceStatusConnecting = true
in.relayTransportConnected = true
in.iceStatusConnected = true
},
want: guard.ConnStatusConnected,
},
{
name: "ICE connected, peer does NOT use relay",
name: "ICE connected, peer does NOT use relay, shared transport down",
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false
in.relayConnected = false
in.iceStatusConnecting = true
in.relayTransportConnected = false
in.iceStatusConnected = true
},
// A peer that does not rely on relay is unaffected by the shared transport:
// relayOK is true, so the first arm matches before the transport is considered.
want: guard.ConnStatusConnected,
},
{
name: "ICE InProgress only, peer does NOT use relay",
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false
in.iceStatusConnecting = false
in.iceStatusConnected = false
in.iceInProgress = true
},
want: guard.ConnStatusConnected,
@@ -150,7 +169,8 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true
in.relayConnected = true
in.iceStatusConnecting = false
in.relayTransportConnected = true
in.iceStatusConnected = false
in.iceInProgress = false
},
want: guard.ConnStatusPartiallyConnected,
@@ -160,21 +180,60 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false
in.relayConnected = false
in.iceStatusConnecting = false
in.iceStatusConnected = false
in.iceInProgress = false
},
want: guard.ConnStatusDisconnected,
},
{
name: "ICE up, peer uses relay but relay down -> partial (relay required, ICE ignored)",
name: "ICE connected, relay down for this peer but the shared transport is up -> disconnected",
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true
in.relayConnected = false
in.iceStatusConnecting = true
in.relayTransportConnected = true
in.iceStatusConnected = true
},
// The transport is fine, so the peer itself is unreachable over relay: it may have
// moved to another server, and only an offer carries its new relay address.
want: guard.ConnStatusDisconnected,
},
{
name: "ICE connected, the shared relay transport is down -> partial",
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true
in.relayConnected = false
in.relayTransportConnected = false
in.iceStatusConnected = true
},
// ICE carries the traffic and the relay transport is restored by the relay client's
// own guard, not by offers, so this must not trigger the aggressive retry.
want: guard.ConnStatusPartiallyConnected,
},
{
name: "ICE only negotiating while the shared relay transport is down -> disconnected",
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true
in.relayConnected = false
in.relayTransportConnected = false
in.iceStatusConnected = false
in.iceInProgress = true
},
// A negotiation in flight is not a working transport, so this peer has no path at
// all and must keep the aggressive retry. Calling it partially connected spends the
// ICE retry budget and parks the guard on the hourly ticker, and nothing wakes it
// when the negotiation then fails: onICEStateDisconnected is only reached once ICE
// has reached Connected (worker_ice.go onConnectionStateChange).
want: guard.ConnStatusDisconnected,
},
{
name: "ICE down and the shared relay transport is down -> disconnected",
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true
in.relayConnected = false
in.relayTransportConnected = false
in.iceStatusConnected = false
in.iceInProgress = false
},
// relayOK = false (peer uses relay but it's down), iceUp = true
// first switch arm fails (relayOK false), relayUsedAndUp = false (relay down),
// falls into default: Disconnected.
want: guard.ConnStatusDisconnected,
},
{
@@ -182,7 +241,7 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false
in.relayConnected = true // not actually used since peer doesn't rely on it
in.iceStatusConnecting = false
in.iceStatusConnected = false
in.iceInProgress = false
},
want: guard.ConnStatusDisconnected,
+5 -3
View File
@@ -14,7 +14,8 @@ type ConnStatus int
const (
// ConnStatusDisconnected means neither ICE nor Relay is connected.
ConnStatusDisconnected ConnStatus = iota
// ConnStatusPartiallyConnected means Relay is connected but ICE is not.
// ConnStatusPartiallyConnected means one transport is usable and the other is not:
// relay connected with ICE down, or ICE connected with the shared relay transport down.
ConnStatusPartiallyConnected
// ConnStatusConnected means all required connections are established.
ConnStatusConnected
@@ -87,8 +88,9 @@ func (g *Guard) SetICEConnDisconnected() {
// - Connected: no action, the peer is fully reachable.
// - Disconnected (neither ICE nor Relay): retries aggressively with exponential backoff (800ms doubling
// up to timeout), never gives up. This ensures rapid recovery when the peer has no connectivity at all.
// - PartiallyConnected (Relay up, ICE not): retries up to 3 times with exponential backoff, then switches
// to one attempt per hour. This limits signaling traffic when relay already provides connectivity.
// - PartiallyConnected (one transport usable, the other not): retries up to 3 times
// with exponential backoff, then switches to one attempt per hour. This limits
// signaling traffic while the peer still has a working path.
//
// External events (relay/ICE disconnect, signal/relay reconnect, candidate changes) reset the retry
// counter and backoff ticker, giving ICE a fresh chance after network conditions change.
+52 -8
View File
@@ -12,6 +12,8 @@ type notifier struct {
serverStateLock sync.Mutex
listenersLock sync.Mutex
listener Listener
peerListWake chan struct{}
peerListStop chan struct{}
currentClientState bool
lastNotification ClientState
lastNumberOfPeers int
@@ -62,7 +64,6 @@ func (n *notifier) setNetworkAvailable(available bool) {
func (n *notifier) setListener(listener Listener) {
n.serverStateLock.Lock()
lastNotification := n.effectiveState(n.lastNotification)
numOfPeers := n.lastNumberOfPeers
fqdnAddress := n.lastFqdnAddress
address := n.lastIPAddress
n.serverStateLock.Unlock()
@@ -70,17 +71,19 @@ func (n *notifier) setListener(listener Listener) {
n.listenersLock.Lock()
defer n.listenersLock.Unlock()
n.stopPeerListDelivererLocked()
n.listener = listener
listener.OnAddressChanged(fqdnAddress, address)
notifyListener(listener, lastNotification)
// run on go routine to avoid on Java layer to call go functions on same thread
go listener.OnPeersListChanged(numOfPeers)
n.startPeerListDelivererLocked(listener)
n.wakePeerListDelivererLocked()
}
func (n *notifier) removeListener() {
n.listenersLock.Lock()
defer n.listenersLock.Unlock()
n.stopPeerListDelivererLocked()
n.listener = nil
}
@@ -178,15 +181,56 @@ func (n *notifier) peerListChanged(numOfPeers int) {
n.serverStateLock.Unlock()
n.listenersLock.Lock()
listener := n.listener
n.listenersLock.Unlock()
defer n.listenersLock.Unlock()
n.wakePeerListDelivererLocked()
}
if listener == nil {
func (n *notifier) startPeerListDelivererLocked(listener Listener) {
wake := make(chan struct{}, 1)
stop := make(chan struct{})
n.peerListWake = wake
n.peerListStop = stop
go n.deliverPeerListChanges(listener, wake, stop)
}
func (n *notifier) stopPeerListDelivererLocked() {
if n.peerListStop == nil {
return
}
close(n.peerListStop)
n.peerListStop = nil
n.peerListWake = nil
}
// run on go routine to avoid on Java layer to call go functions on same thread
go listener.OnPeersListChanged(numOfPeers)
func (n *notifier) wakePeerListDelivererLocked() {
if n.peerListWake == nil {
return
}
select {
case n.peerListWake <- struct{}{}:
default:
}
}
func (n *notifier) deliverPeerListChanges(listener Listener, wake <-chan struct{}, stop <-chan struct{}) {
for {
select {
case <-stop:
return
case <-wake:
}
select {
case <-stop:
return
default:
}
n.serverStateLock.Lock()
numOfPeers := n.lastNumberOfPeers
n.serverStateLock.Unlock()
listener.OnPeersListChanged(numOfPeers)
}
}
func (n *notifier) localAddressChanged(fqdn, address string) {
+155
View File
@@ -2,7 +2,9 @@ package peer
import (
"sync"
"sync/atomic"
"testing"
"time"
)
type mocListener struct {
@@ -115,3 +117,156 @@ func Test_notifier_RemoveListener(t *testing.T) {
t.Errorf("invalid state: %d", listener.peers)
}
}
type coalescingListener struct {
final int
calls atomic.Int32
inFlight atomic.Int32
maxInFlight atomic.Int32
last atomic.Int32
done chan struct{}
entered chan struct{}
release chan struct{}
once sync.Once
}
func (l *coalescingListener) OnStateChanged(ClientState) {}
func (l *coalescingListener) OnConnected() {}
func (l *coalescingListener) OnDisconnected() {}
func (l *coalescingListener) OnConnecting() {}
func (l *coalescingListener) OnDisconnecting() {}
func (l *coalescingListener) OnAddressChanged(string, string) {}
func (l *coalescingListener) OnPeersListChanged(size int) {
current := l.inFlight.Add(1)
for {
seen := l.maxInFlight.Load()
if current <= seen || l.maxInFlight.CompareAndSwap(seen, current) {
break
}
}
if l.calls.Add(1) == 1 && l.entered != nil {
close(l.entered)
}
if l.release != nil {
<-l.release
}
time.Sleep(time.Millisecond)
l.last.Store(int32(size))
l.inFlight.Add(-1)
if size == l.final {
l.once.Do(func() { close(l.done) })
}
}
func Test_notifier_PeerListChangedCoalesces(t *testing.T) {
const events = 1000
listener := &coalescingListener{final: events, done: make(chan struct{})}
n := newNotifier()
n.setListener(listener)
for i := 1; i <= events; i++ {
n.peerListChanged(i)
}
select {
case <-listener.done:
case <-time.After(5 * time.Second):
t.Fatalf("last peer count not delivered, last seen: %d", listener.last.Load())
}
if got := listener.maxInFlight.Load(); got != 1 {
t.Errorf("concurrent deliveries: %d, expected 1", got)
}
if got := listener.calls.Load(); got >= events {
t.Errorf("deliveries not coalesced: %d calls for %d events", got, events)
}
}
func Test_notifier_SetListenerStopsPreviousDeliverer(t *testing.T) {
old := &coalescingListener{final: -1}
replacement := &coalescingListener{final: 7, done: make(chan struct{})}
n := newNotifier()
n.setListener(old)
oldStop := n.peerListStop
n.peerListChanged(7)
n.setListener(replacement)
select {
case <-oldStop:
default:
t.Fatal("old deliverer not stopped on listener replacement")
}
waitFor(t, replacement.done, "replacement listener not notified")
}
func Test_notifier_RemoveListenerStopsDeliverer(t *testing.T) {
n := newNotifier()
n.setListener(&coalescingListener{final: -1})
stop := n.peerListStop
n.removeListener()
select {
case <-stop:
default:
t.Fatal("deliverer not stopped on listener removal")
}
}
func Test_notifier_DelivererExitsAfterInFlightCallback(t *testing.T) {
listener := &coalescingListener{
final: -1,
entered: make(chan struct{}),
release: make(chan struct{}),
}
n := newNotifier()
wake := make(chan struct{}, 1)
stop := make(chan struct{})
exited := make(chan struct{})
go func() {
n.deliverPeerListChanges(listener, wake, stop)
close(exited)
}()
wake <- struct{}{}
waitFor(t, listener.entered, "listener not called")
n.peerListChanged(7)
wake <- struct{}{}
close(stop)
close(listener.release)
waitFor(t, exited, "deliverer did not exit after stop")
if got := listener.calls.Load(); got != 1 {
t.Errorf("deliverer ran %d callbacks after stop, expected only the in-flight one", got)
}
if got := listener.last.Load(); got == 7 {
t.Errorf("deliverer delivered the peer count queued after stop")
}
}
func Test_notifier_DelivererPrefersStopOverPendingWake(t *testing.T) {
listener := &coalescingListener{final: -1}
n := newNotifier()
wake := make(chan struct{}, 1)
stop := make(chan struct{})
wake <- struct{}{}
close(stop)
n.deliverPeerListChanges(listener, wake, stop)
if got := listener.calls.Load(); got != 0 {
t.Errorf("deliverer ran %d callbacks with stop closed, expected 0", got)
}
}
func waitFor(t *testing.T, ch <-chan struct{}, msg string) {
t.Helper()
select {
case <-ch:
case <-time.After(5 * time.Second):
t.Fatal(msg)
}
}
-35
View File
@@ -18,9 +18,7 @@ import (
"google.golang.org/protobuf/types/known/durationpb"
"google.golang.org/protobuf/types/known/timestamppb"
firewall "github.com/netbirdio/netbird/client/firewall/manager"
"github.com/netbirdio/netbird/client/iface/configurer"
"github.com/netbirdio/netbird/client/internal/ingressgw"
"github.com/netbirdio/netbird/client/internal/relay"
"github.com/netbirdio/netbird/client/proto"
"github.com/netbirdio/netbird/route"
@@ -161,7 +159,6 @@ type FullStatus struct {
RosenpassState RosenpassState
Relays []relay.ProbeResult
NSGroupStates []NSGroupState
NumOfForwardingRules int
LazyConnectionEnabled bool
Events []*proto.SystemEvent
}
@@ -247,8 +244,6 @@ type Status struct {
// read it without taking mux.
networksRevision atomic.Uint64
ingressGwMgr *ingressgw.Manager
routeIDLookup routeIDLookup
wgIface WGIfaceStatus
}
@@ -276,12 +271,6 @@ func (d *Status) SetRelayMgr(manager *relayClient.Manager) {
d.relayMgr = manager
}
func (d *Status) SetIngressGwMgr(ingressGwMgr *ingressgw.Manager) {
d.mux.Lock()
defer d.mux.Unlock()
d.ingressGwMgr = ingressGwMgr
}
// ReplaceOfflinePeers replaces
func (d *Status) ReplaceOfflinePeers(replacement []State) {
d.mux.Lock()
@@ -332,18 +321,6 @@ func (d *Status) GetPeer(peerPubKey string) (State, error) {
return state, nil
}
func (d *Status) PeerByIP(ip string) (string, bool) {
d.mux.RLock()
defer d.mux.RUnlock()
for _, state := range d.peers {
if state.IP == ip {
return state.FQDN, true
}
}
return "", false
}
// PeerStateByIP returns the full peer State for the given tunnel IP.
// Matches against either the IPv4 (State.IP) or IPv6 (State.IPv6) tunnel
// address so dual-stack peers are reachable on either family. Only
@@ -1163,16 +1140,6 @@ func (d *Status) GetRelayStates() []relay.ProbeResult {
return relayStates
}
func (d *Status) ForwardingRules() []firewall.ForwardRule {
d.mux.RLock()
defer d.mux.RUnlock()
if d.ingressGwMgr == nil {
return nil
}
return d.ingressGwMgr.Rules()
}
func (d *Status) GetDNSStates() []NSGroupState {
d.mux.RLock()
defer d.mux.RUnlock()
@@ -1207,7 +1174,6 @@ func (d *Status) GetFullStatus() FullStatus {
Relays: d.GetRelayStates(),
RosenpassState: d.GetRosenpassState(),
NSGroupStates: d.GetDNSStates(),
NumOfForwardingRules: len(d.ForwardingRules()),
LazyConnectionEnabled: d.GetLazyConnection(),
}
@@ -1579,7 +1545,6 @@ func (fs FullStatus) ToProto() *proto.FullStatus {
pbFullStatus.LocalPeerState.WgPort = int32(fs.LocalPeerState.WgPort)
pbFullStatus.LocalPeerState.RosenpassPermissive = fs.RosenpassState.Permissive
pbFullStatus.LocalPeerState.RosenpassEnabled = fs.RosenpassState.Enabled
pbFullStatus.NumberOfForwardingRules = int32(fs.NumOfForwardingRules)
pbFullStatus.LazyConnectionEnabled = fs.LazyConnectionEnabled
pbFullStatus.LocalPeerState.Networks = maps.Keys(fs.LocalPeerState.Routes)
+4
View File
@@ -101,6 +101,10 @@ func (w *WorkerRelay) RelayIsSupportedLocally() bool {
return w.relayManager.HasRelayAddress()
}
func (w *WorkerRelay) IsTransportConnected() bool {
return w.relayManager.Ready()
}
func (w *WorkerRelay) CloseConn() {
w.relayLock.Lock()
conn := w.relayedConn
+377 -106
View File
@@ -10,7 +10,6 @@ import (
"os"
"os/user"
"path/filepath"
"reflect"
"runtime"
"slices"
"strings"
@@ -198,6 +197,11 @@ type Config struct {
MTU uint16
// probing marks a config that exists only to be compared against and then
// thrown away, so apply() can skip the work that feeds no verdict.
// Unexported, so it never reaches the JSON.
probing bool
// policy is the MDM policy that produced the currently-set values
// for any MDM-enforced fields. Set by ApplyMDMPolicy on every
// invocation. Never persisted to disk. Callers query enforcement
@@ -300,9 +304,11 @@ func fileExists(path string) (bool, error) {
return false, err
}
// createNewConfig creates a new config generating a new Wireguard key and saving to file
func createNewConfig(input ConfigInput) (*Config, error) {
config := &Config{
// newConfigSkeleton returns the field values a brand-new profile config starts
// from, before apply() fills in the rest. Shared with the dry-run baseline so
// the two cannot disagree about what "a new config" means.
func newConfigSkeleton() *Config {
return &Config{
// defaults to false only for new (post 0.26) configurations
ServerSSHAllowed: util.False(),
// Remote jobs are an explicit opt-in and default off, including for
@@ -310,6 +316,91 @@ func createNewConfig(input ConfigInput) (*Config, error) {
RemoteJobsAllowed: util.False(),
WgPort: iface.DefaultWgPort,
}
}
// resolveUnsetDefaults is the single place where an optional field that carries
// no value gets one, and the only place that states what each of those defaults
// is. apply() runs it before it compares anything, and that ordering is the
// point: with the values named, every comparison below it diffs values instead
// of presence.
//
// Presence-based comparison is what broke `netbird up` for a client configured
// through the environment. These fields mean "the effective default" when they
// hold nothing — every consumer already reads a nil as the value resolved here,
// the SSH toggles in engine_ssh.go and the network monitor in
// createEngineConfig — so naming them changes nothing about what runs. But
// while they stayed nil, an input restating the default read as a change, and
// since the CLI sends every flag whose value came from an environment variable
// on each `netbird up`, a client with NB_ENABLE_SSH_ROOT=false restated it
// every time and the update-settings gate refused it.
//
// Filling a field in is not a settings change, so a caller measuring change
// must not read the returned bool as one: see WouldChange, which runs a pass
// for this and discards its verdict.
//
// ServerSSHAllowed is the one field whose default depends on the config's age.
// A brand-new profile gets false from newConfigSkeleton, which runs before
// this, so what is resolved here is only the legacy case: a config written by a
// version that had no such field keeps SSH on, for backwards compatibility.
func (config *Config) resolveUnsetDefaults() (updated bool) {
// Fields that default to false on every platform.
for _, field := range []**bool{
&config.EnableSSHRoot,
&config.EnableSSHSFTP,
&config.EnableSSHLocalPortForwarding,
&config.EnableSSHRemotePortForwarding,
&config.DisableSSHAuth,
// Remote jobs are an explicit opt-in: unlike SSH, a pre-existing config
// with no value defaults to disabled rather than being turned on.
&config.RemoteJobsAllowed,
} {
if *field == nil {
*field = util.False()
updated = true
}
}
if config.DisableNotifications == nil {
log.Infof("setting notifications to disabled by default")
config.DisableNotifications = util.True()
updated = true
}
if config.SSHJWTCacheTTL == nil {
// A zero TTL disables the JWT cache, which is what no value meant.
config.SSHJWTCacheTTL = new(int)
updated = true
}
if config.NetworkMonitor == nil {
// network monitoring is on by default on windows and darwin clients
enabled := runtime.GOOS == "windows" || runtime.GOOS == "darwin"
config.NetworkMonitor = &enabled
updated = true
}
if config.ServerSSHAllowed == nil {
if runtime.GOOS == "android" {
// default to disabled SSH on Android for security
log.Infof("setting SSH server to false by default on Android")
config.ServerSSHAllowed = util.False()
} else {
// enables SSH for configs from old versions to preserve backwards compatibility
log.Infof("falling back to enabled SSH server for pre-existing configuration")
config.ServerSSHAllowed = util.True()
}
updated = true
}
return updated
}
// createNewConfig resolves a new config in memory, with no identity: whoever
// needs the peer's keys calls EnsureIdentity and persists the result, so a read
// that lands on a missing file cannot hand back a config carrying keys that
// nothing will ever write down.
func createNewConfig(input ConfigInput) (*Config, error) {
config := newConfigSkeleton()
if _, err := config.apply(input); err != nil {
return nil, err
@@ -318,6 +409,52 @@ func createNewConfig(input ConfigInput) (*Config, error) {
return config, nil
}
// createProvisionedConfig is createNewConfig plus the peer's identity, for the
// callers that go on to persist the config or to connect with it.
func createProvisionedConfig(input ConfigInput) (*Config, error) {
config, err := createNewConfig(input)
if err != nil {
return nil, err
}
if _, err := config.EnsureIdentity(); err != nil {
return nil, err
}
return config, nil
}
// EnsureIdentity generates the keys that identify this peer if the config does
// not carry them yet, reporting whether it had to generate any.
//
// It is deliberately not part of apply(). Everything apply() fills in is a
// default it can recompute on the next read, but a generated key is not: it
// has to be persisted, or the peer comes back with a different WireGuard
// identity and re-registers. Having apply() generate keys is what forced every
// read of a config to write it back — so identity provisioning is its own step
// now, and the callers that perform it write the result out explicitly.
func (config *Config) EnsureIdentity() (bool, error) {
generated := false
if config.PrivateKey == "" {
log.Infof("generated new Wireguard key")
config.PrivateKey = generateKey()
generated = true
}
if config.SSHKey == "" {
log.Infof("generated new SSH key")
pem, err := ssh.GeneratePrivateKey(ssh.ED25519)
if err != nil {
return generated, err
}
config.SSHKey = string(pem)
generated = true
}
return generated, nil
}
func (config *Config) apply(input ConfigInput) (updated bool, err error) {
if config.Name != "" {
sanitized, err := sanitizeDisplayName(config.Name)
@@ -329,6 +466,13 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true
}
}
// Every optional field gets its value here, before anything below compares
// one. See resolveUnsetDefaults for why that ordering is the point.
if config.resolveUnsetDefaults() {
updated = true
}
if config.ManagementURL == nil {
log.Infof("using default Management URL %s", DefaultManagementURL)
config.ManagementURL, err = parseURL("Management URL", DefaultManagementURL)
@@ -336,20 +480,21 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
return false, err
}
}
if input.ManagementURL != "" && input.ManagementURL != config.ManagementURL.String() {
log.Infof("new Management URL provided, updated to %#v (old value %#v)",
input.ManagementURL, config.ManagementURL.String())
// The comparison is on the endpoint the URL addresses, not on its
// spelling: the same endpoint can be written several ways (an implicit
// :443, a trailing slash, a different host case), and treating an
// equivalent URL as new would rewrite the config and report a settings
// change where the configuration does not actually change.
if input.ManagementURL != "" {
URL, err := parseURL("Management URL", input.ManagementURL)
if err != nil {
return false, err
}
config.ManagementURL = URL
updated = true
} else if config.ManagementURL == nil {
log.Infof("using default Management URL %s", DefaultManagementURL)
config.ManagementURL, err = parseURL("Management URL", DefaultManagementURL)
if err != nil {
return false, err
if !SameServiceURL(URL, config.ManagementURL) {
log.Infof("new Management URL provided, updated to %#v (old value %#v)",
URL.String(), config.ManagementURL.String())
config.ManagementURL = URL
updated = true
}
}
@@ -360,31 +505,20 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
return false, err
}
}
if input.AdminURL != "" && input.AdminURL != config.AdminURL.String() {
log.Infof("new Admin Panel URL provided, updated to %#v (old value %#v)",
input.AdminURL, config.AdminURL.String())
// The admin panel is opened, not dialed, so unlike the Management URL its
// path is part of what identifies it: a panel served under /netbird is not
// the one served at the root.
if input.AdminURL != "" {
newURL, err := parseURL("Admin Panel URL", input.AdminURL)
if err != nil {
return updated, err
}
config.AdminURL = newURL
updated = true
}
if config.PrivateKey == "" {
log.Infof("generated new Wireguard key")
config.PrivateKey = generateKey()
updated = true
}
if config.SSHKey == "" {
log.Infof("generated new SSH key")
pem, err := ssh.GeneratePrivateKey(ssh.ED25519)
if err != nil {
return false, err
if !SameServiceURLIncludingPath(newURL, config.AdminURL) {
log.Infof("new Admin Panel URL provided, updated to %#v (old value %#v)",
newURL.String(), config.AdminURL.String())
config.AdminURL = newURL
updated = true
}
config.SSHKey = string(pem)
updated = true
}
if input.WireguardPort != nil && *input.WireguardPort != config.WgPort {
@@ -405,7 +539,14 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true
}
if input.NATExternalIPs != nil && !reflect.DeepEqual(config.NATExternalIPs, input.NATExternalIPs) {
// slices.Equal, not reflect.DeepEqual, and for the same reason the DNS
// labels below use it: DeepEqual calls a nil slice and an empty one
// different, while both mean "no NAT mappings". A profile stores the
// absent list as JSON null and reads it back nil, and `netbird up` sends
// CleanNATExternalIPs — an empty list — whenever NB_EXTERNAL_IP_MAP is set
// to nothing, so the two met on every start and the gate read a no-op as a
// settings change.
if input.NATExternalIPs != nil && !slices.Equal(config.NATExternalIPs, input.NATExternalIPs) {
log.Infof("updating NAT External IP [ %s ] (old value: [ %s ])",
strings.Join(input.NATExternalIPs, " "),
strings.Join(config.NATExternalIPs, " "))
@@ -443,21 +584,12 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true
}
if input.NetworkMonitor != nil && (config.NetworkMonitor == nil || *input.NetworkMonitor != *config.NetworkMonitor) {
if input.NetworkMonitor != nil && *input.NetworkMonitor != *config.NetworkMonitor {
log.Infof("switching Network Monitor to %t", *input.NetworkMonitor)
config.NetworkMonitor = input.NetworkMonitor
updated = true
}
if config.NetworkMonitor == nil {
// enable network monitoring by default on windows and darwin clients
if runtime.GOOS == "windows" || runtime.GOOS == "darwin" {
enabled := true
config.NetworkMonitor = &enabled
updated = true
}
}
if input.CustomDNSAddress != nil && string(input.CustomDNSAddress) != config.CustomDNSAddress {
log.Infof("updating custom DNS address %#v (old value %#v)",
string(input.CustomDNSAddress), config.CustomDNSAddress)
@@ -490,7 +622,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true
}
if input.ServerSSHAllowed != nil && (config.ServerSSHAllowed == nil || *input.ServerSSHAllowed != *config.ServerSSHAllowed) {
if input.ServerSSHAllowed != nil && *input.ServerSSHAllowed != *config.ServerSSHAllowed {
if *input.ServerSSHAllowed {
log.Infof("enabling SSH server")
} else {
@@ -498,20 +630,9 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
}
config.ServerSSHAllowed = input.ServerSSHAllowed
updated = true
} else if config.ServerSSHAllowed == nil {
if runtime.GOOS == "android" {
// default to disabled SSH on Android for security
log.Infof("setting SSH server to false by default on Android")
config.ServerSSHAllowed = util.False()
} else {
// enables SSH for configs from old versions to preserve backwards compatibility
log.Infof("falling back to enabled SSH server for pre-existing configuration")
config.ServerSSHAllowed = util.True()
}
updated = true
}
if input.RemoteJobsAllowed != nil && (config.RemoteJobsAllowed == nil || *input.RemoteJobsAllowed != *config.RemoteJobsAllowed) {
if input.RemoteJobsAllowed != nil && *input.RemoteJobsAllowed != *config.RemoteJobsAllowed {
if *input.RemoteJobsAllowed {
log.Infof("enabling remote jobs")
} else {
@@ -519,14 +640,9 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
}
config.RemoteJobsAllowed = input.RemoteJobsAllowed
updated = true
} else if config.RemoteJobsAllowed == nil {
// Remote jobs are an explicit opt-in: unlike SSH, a pre-existing config
// with no value defaults to disabled rather than being turned on.
config.RemoteJobsAllowed = util.False()
updated = true
}
if input.EnableSSHRoot != nil && (config.EnableSSHRoot == nil || *input.EnableSSHRoot != *config.EnableSSHRoot) {
if input.EnableSSHRoot != nil && *input.EnableSSHRoot != *config.EnableSSHRoot {
if *input.EnableSSHRoot {
log.Infof("enabling SSH root login")
} else {
@@ -536,7 +652,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true
}
if input.EnableSSHSFTP != nil && (config.EnableSSHSFTP == nil || *input.EnableSSHSFTP != *config.EnableSSHSFTP) {
if input.EnableSSHSFTP != nil && *input.EnableSSHSFTP != *config.EnableSSHSFTP {
if *input.EnableSSHSFTP {
log.Infof("enabling SSH SFTP subsystem")
} else {
@@ -546,7 +662,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true
}
if input.EnableSSHLocalPortForwarding != nil && (config.EnableSSHLocalPortForwarding == nil || *input.EnableSSHLocalPortForwarding != *config.EnableSSHLocalPortForwarding) {
if input.EnableSSHLocalPortForwarding != nil && *input.EnableSSHLocalPortForwarding != *config.EnableSSHLocalPortForwarding {
if *input.EnableSSHLocalPortForwarding {
log.Infof("enabling SSH local port forwarding")
} else {
@@ -556,7 +672,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true
}
if input.EnableSSHRemotePortForwarding != nil && (config.EnableSSHRemotePortForwarding == nil || *input.EnableSSHRemotePortForwarding != *config.EnableSSHRemotePortForwarding) {
if input.EnableSSHRemotePortForwarding != nil && *input.EnableSSHRemotePortForwarding != *config.EnableSSHRemotePortForwarding {
if *input.EnableSSHRemotePortForwarding {
log.Infof("enabling SSH remote port forwarding")
} else {
@@ -566,7 +682,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true
}
if input.DisableSSHAuth != nil && (config.DisableSSHAuth == nil || *input.DisableSSHAuth != *config.DisableSSHAuth) {
if input.DisableSSHAuth != nil && *input.DisableSSHAuth != *config.DisableSSHAuth {
if *input.DisableSSHAuth {
log.Infof("disabling SSH authentication")
} else {
@@ -576,7 +692,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true
}
if input.SSHJWTCacheTTL != nil && (config.SSHJWTCacheTTL == nil || *input.SSHJWTCacheTTL != *config.SSHJWTCacheTTL) {
if input.SSHJWTCacheTTL != nil && *input.SSHJWTCacheTTL != *config.SSHJWTCacheTTL {
log.Infof("updating SSH JWT cache TTL to %d seconds", *input.SSHJWTCacheTTL)
config.SSHJWTCacheTTL = input.SSHJWTCacheTTL
updated = true
@@ -659,13 +775,16 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true
}
if input.SyncMessageVersion != nil && *input.SyncMessageVersion != *config.SyncMessageVersion {
// Assigning the pointer, not writing through it: a config that carries no
// version yet would otherwise be a nil dereference, and a panic inside a
// request handler is not a way to fail.
if input.SyncMessageVersion != nil && (config.SyncMessageVersion == nil || *input.SyncMessageVersion != *config.SyncMessageVersion) {
log.Infof("setting SyncMessageVersion to %v", *input.SyncMessageVersion)
*config.SyncMessageVersion = *input.SyncMessageVersion
config.SyncMessageVersion = input.SyncMessageVersion
updated = true
}
if input.DisableNotifications != nil && (config.DisableNotifications == nil || *input.DisableNotifications != *config.DisableNotifications) {
if input.DisableNotifications != nil && *input.DisableNotifications != *config.DisableNotifications {
if *input.DisableNotifications {
log.Infof("disabling notifications")
} else {
@@ -675,24 +794,24 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true
}
if config.DisableNotifications == nil {
disabled := true
config.DisableNotifications = &disabled
log.Infof("setting notifications to disabled by default")
updated = true
}
if input.ClientCertKeyPath != "" {
// Compared, not just assigned: restating the path a config already holds
// changes nothing, and reporting it as an update makes a caller that
// re-sends its own configuration look like one asking to change it.
if input.ClientCertKeyPath != "" && input.ClientCertKeyPath != config.ClientCertKeyPath {
config.ClientCertKeyPath = input.ClientCertKeyPath
updated = true
}
if input.ClientCertPath != "" {
if input.ClientCertPath != "" && input.ClientCertPath != config.ClientCertPath {
config.ClientCertPath = input.ClientCertPath
updated = true
}
if config.ClientCertPath != "" && config.ClientCertKeyPath != "" {
// Not on a probe: the loaded pair feeds the connection, never the
// comparison, and this would otherwise run on every gated SetConfig and
// Login — twice per request — including those that are refused or change
// nothing, logging an error per request when the files are missing.
if !config.probing && config.ClientCertPath != "" && config.ClientCertKeyPath != "" {
cert, err := tls.LoadX509KeyPair(config.ClientCertPath, config.ClientCertKeyPath)
if err != nil {
log.Error("Failed to load mTLS cert/key pair: ", err)
@@ -886,6 +1005,49 @@ func ParseServiceURL(serviceName, serviceURL string) (*url.URL, error) {
return parseURL(serviceName, serviceURL)
}
// SameServiceURL reports whether two service URLs address the same endpoint:
// same scheme, same host compared case-insensitively as DNS names are, and
// same effective port, where an absent port means the scheme's default.
//
// This is the one comparison every caller deciding "did this URL change?" must
// use. A string comparison answers a different question: "https://host",
// "https://host/" and "https://HOST:443" are one endpoint written three ways,
// and reading them as three values makes a client that restates its own
// management URL look like a client asking to be repointed. A nil operand
// matches only another nil one.
//
// The path plays no part: a management URL is dialed, and only its host and
// port are. util.SameServiceURL is this comparison plus the path, which is
// what SameServiceURLIncludingPath needs and delegates to.
func SameServiceURL(a, b *url.URL) bool {
if a == nil || b == nil {
return a == b
}
return strings.EqualFold(a.Scheme, b.Scheme) &&
strings.EqualFold(a.Hostname(), b.Hostname()) &&
util.ServiceURLPort(a) == util.ServiceURLPort(b)
}
// SameServiceURLIncludingPath is SameServiceURL plus everything a URL carries
// past its endpoint: path, query, fragment and userinfo.
//
// Use it for a URL that gets opened rather than dialed. The admin panel can
// live under a path, so two URLs with the same endpoint and different paths are
// two different panels — where for a URL the client dials over gRPC only the
// endpoint is ever used. Equivalent spellings still compare equal: a missing
// path and "/" are the same root, and so is a trailing slash on any path.
func SameServiceURLIncludingPath(a, b *url.URL) bool {
if a == nil || b == nil {
return a == b
}
return util.SameServiceURL(a, b) &&
a.RawQuery == b.RawQuery &&
a.Fragment == b.Fragment &&
a.User.String() == b.User.String()
}
func parseURL(serviceName, serviceURL string) (*url.URL, error) {
parsedMgmtURL, err := url.ParseRequestURI(serviceURL)
if err != nil {
@@ -930,6 +1092,84 @@ func isPreSharedKeyHidden(preSharedKey *string) bool {
return false
}
// WouldChange reports whether applying input would modify any field the
// config persists, leaving the receiver untouched. It is the dry-run half of
// UpdateConfig and reuses the very same diff logic (Config.apply), so a
// caller asking "is this a settings change?" cannot drift from what an
// actual update would do, nor go stale when a new field is added.
//
// A redacted pre-shared key is collapsed to "unset" exactly as
// UpdateOrCreateConfig does, so a UI that round-trips the mask is not read as
// a request for a new key.
//
// A nil receiver means the profile holds no config yet, so the baseline is the
// config the daemon would create for it: input values matching those defaults
// change nothing, anything else does.
func (config *Config) WouldChange(input ConfigInput) (bool, error) {
probe := config.clone()
if probe == nil {
baseline, err := newDryRunBaseline(input.ConfigPath)
if err != nil {
return true, fmt.Errorf("build default config baseline: %w", err)
}
probe = baseline
}
probe.probing = true
// Normalize before measuring. apply() reports two different things through
// one bool: an input that changed a value, and a field it had to fill in
// because the config carried none. Only the first is a settings change, so
// the filling-in gets a pass of its own whose verdict is discarded, and the
// pass that answers the caller runs against a config with nothing left to
// fill in.
//
// Readers already hand out normalized configs — readConfig applies an empty
// input for this very reason — so this is normally a no-op. But a gate that
// refuses a request must not depend on where its caller got the config
// from, and it must not start reading "this profile predates a field" as
// "the caller asked for a change" the day someone adds one.
if _, err := probe.apply(ConfigInput{ConfigPath: input.ConfigPath}); err != nil {
return true, fmt.Errorf("normalize the config to diff against: %w", err)
}
if isPreSharedKeyHidden(input.PreSharedKey) {
input.PreSharedKey = nil
}
return probe.apply(input)
}
// newDryRunBaseline builds the config a brand-new profile would start from, for
// a dry run to compare an input against. It is createNewConfig without the
// identity: this config exists only to be compared against and thrown away, and
// no ConfigInput field maps to either key.
func newDryRunBaseline(configPath string) (*Config, error) {
baseline := newConfigSkeleton()
if _, err := baseline.apply(ConfigInput{ConfigPath: configPath}); err != nil {
return nil, err
}
return baseline, nil
}
// clone returns a copy of the config that apply can be run against without the
// original observing the writes, or nil for a nil receiver. Only what apply
// mutates in place needs detaching, which is the slices it replaces or appends
// to: every pointer field it touches is reassigned rather than written through,
// and ClientCertKeyPair is only overwritten.
func (config *Config) clone() *Config {
if config == nil {
return nil
}
probe := *config
probe.IFaceBlackList = slices.Clone(config.IFaceBlackList)
probe.NATExternalIPs = slices.Clone(config.NATExternalIPs)
probe.DNSLabels = slices.Clone(config.DNSLabels)
return &probe
}
// UpdateConfig update existing configuration according to input configuration and return with the configuration
func UpdateConfig(input ConfigInput) (*Config, error) {
configExists, err := fileExists(input.ConfigPath)
@@ -940,6 +1180,14 @@ func UpdateConfig(input ConfigInput) (*Config, error) {
return nil, fmt.Errorf("config file %s does not exist", input.ConfigPath)
}
// A UI that round-trips the mask GetConfig hands it back is asking to keep
// the stored key, not to set the mask as the new one. UpdateOrCreateConfig
// and DirectUpdateOrCreateConfig already collapse it; this one did not, so
// the same round-trip through SetConfig replaced the key with asterisks.
if isPreSharedKeyHidden(input.PreSharedKey) {
input.PreSharedKey = nil
}
return update(input)
}
@@ -951,7 +1199,7 @@ func UpdateOrCreateConfig(input ConfigInput) (*Config, error) {
}
if !configExists {
log.Infof("generating new config %s", input.ConfigPath)
cfg, err := createNewConfig(input)
cfg, err := createProvisionedConfig(input)
if err != nil {
return nil, err
}
@@ -976,12 +1224,20 @@ func update(input ConfigInput) (*Config, error) {
return nil, err
}
// A write path is a provisioning point: a stored profile can legitimately
// carry no identity (a mobile logout clears the keys in place), and the
// next config write is what has to mint a new one. Reads leave that alone.
identityGenerated, err := config.EnsureIdentity()
if err != nil {
return nil, err
}
updated, err := config.apply(input)
if err != nil {
return nil, err
}
if updated {
if updated || identityGenerated {
if err := util.WriteJson(context.Background(), input.ConfigPath, config); err != nil {
return nil, err
}
@@ -990,8 +1246,8 @@ func update(input ConfigInput) (*Config, error) {
return config, nil
}
// GetConfig read config file and return with Config and if it was created. Errors out if it does not exist
func GetConfig(configPath string) (*Config, error) {
// GetExistingConfig reads and returns the config if it exists on disk. Fails otherwise.
func GetExistingConfig(configPath string) (*Config, error) {
return readConfig(configPath, false)
}
@@ -1074,17 +1330,27 @@ func UpdateOldManagementURL(ctx context.Context, config *Config, configPath stri
return newConfig, nil
}
// CreateInMemoryConfig generate a new config but do not write out it to the store
// CreateInMemoryConfig generate a new config but do not write out it to the store.
// It carries an identity: callers connect with what they get back.
func CreateInMemoryConfig(input ConfigInput) (*Config, error) {
return createNewConfig(input)
return createProvisionedConfig(input)
}
// ReadConfig read config file and return with Config. If it is not exists create a new with default values
func ReadConfig(configPath string) (*Config, error) {
// ReadConfigOrDefault reads the profile config at configPath, or resolves the
// default config in memory when the file does not exist. It never writes, and
// never mints an identity — EnsureIdentity is where that happens, so the
// caller that provisions is also the one that persists.
func ReadConfigOrDefault(configPath string) (*Config, error) {
return readConfig(configPath, true)
}
// ReadConfig read config file and return with Config. If it is not exists create a new with default values
// readConfig reads the profile config at configPath. createIfMissing resolves a
// default config in memory when the file is absent, rather than erroring.
//
// Reads are pure. This used to write the config back whenever apply() had to
// fill in a default the file was missing, which quietly made every reader a
// writer: a gate deciding whether to refuse a request, a UI listing profiles,
// a mobile getter reading a single preference.
func readConfig(configPath string, createIfMissing bool) (*Config, error) {
configExists, err := fileExists(configPath)
if err != nil {
@@ -1102,12 +1368,8 @@ func readConfig(configPath string, createIfMissing bool) (*Config, error) {
return nil, err
}
// initialize through apply() without changes
if changed, err := config.apply(ConfigInput{}); err != nil {
if _, err := config.apply(ConfigInput{}); err != nil {
return nil, err
} else if changed {
if err = WriteOutConfig(configPath, config); err != nil {
return nil, err
}
}
return config, nil
@@ -1115,13 +1377,7 @@ func readConfig(configPath string, createIfMissing bool) (*Config, error) {
return nil, fmt.Errorf("config file %s does not exist", configPath)
}
cfg, err := createNewConfig(ConfigInput{ConfigPath: configPath})
if err != nil {
return nil, err
}
err = WriteOutConfig(configPath, cfg)
return cfg, err
return createNewConfig(ConfigInput{ConfigPath: configPath})
}
// WriteOutConfig write put the prepared config to the given path
@@ -1144,7 +1400,7 @@ func DirectUpdateOrCreateConfig(input ConfigInput) (*Config, error) {
}
if !configExists {
log.Infof("generating new config %s", input.ConfigPath)
cfg, err := createNewConfig(input)
cfg, err := createProvisionedConfig(input)
if err != nil {
return nil, err
}
@@ -1171,12 +1427,18 @@ func directUpdate(input ConfigInput) (*Config, error) {
return nil, err
}
// Same provisioning point as update(); see the note there.
identityGenerated, err := config.EnsureIdentity()
if err != nil {
return nil, err
}
updated, err := config.apply(input)
if err != nil {
return nil, err
}
if updated {
if updated || identityGenerated {
if err := util.DirectWriteJson(context.Background(), input.ConfigPath, config); err != nil {
return nil, err
}
@@ -1198,7 +1460,16 @@ func ConfigToJSON(config *Config) (string, error) {
// ConfigFromJSON deserializes a JSON string to a Config struct.
// This is useful for restoring config from alternative storage mechanisms.
// After unmarshaling, defaults are applied to ensure the config is fully initialized.
// After unmarshaling, defaults are applied to ensure the config is fully
// initialized.
//
// The peer identity is deliberately none of its business, in either direction.
// It does not generate one: a read cannot hand back keys that nothing will
// write down (see ReadConfigOrDefault). Nor does it refuse a document that
// carries none, because a config legitimately has no identity between a logout
// and the next login — mobile logout clears both keys in place — and this is
// also the deserializer the iOS SDK copies a config through. Whoever goes on
// to connect is where an absent identity has to be answered.
func ConfigFromJSON(jsonStr string) (*Config, error) {
config := &Config{}
err := json.Unmarshal([]byte(jsonStr), config)
@@ -0,0 +1,44 @@
package profilemanager
import (
"path/filepath"
"testing"
"github.com/stretchr/testify/require"
)
// The serialized form is how the tvOS SDK stores a profile and how the iOS SDK
// copies one in memory, so it must round-trip whatever a profile legitimately
// holds — including no identity at all, which is the state mobile logout leaves
// behind when it clears both keys in place. Refusing that document here broke
// logout, profile switching and the login that follows them.
func TestConfigFromJSONRoundTripsALoggedOutProfile(t *testing.T) {
path := filepath.Join(t.TempDir(), "exported.json")
stored, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path, ManagementURL: DefaultManagementURL})
require.NoError(t, err)
require.NotEmpty(t, stored.PrivateKey, "a provisioned config is the fixture this test starts from")
require.NotEmpty(t, stored.SSHKey)
exported, err := ConfigToJSON(stored)
require.NoError(t, err)
restored, err := ConfigFromJSON(exported)
require.NoError(t, err, "a config exported after a login must load")
require.Equal(t, stored.PrivateKey, restored.PrivateKey, "the restored peer is not the stored one")
require.Equal(t, stored.SSHKey, restored.SSHKey)
// What mobile logout leaves on disk.
loggedOut := stored.clone()
loggedOut.PrivateKey = ""
loggedOut.SSHKey = ""
document, err := ConfigToJSON(loggedOut)
require.NoError(t, err)
reloaded, err := ConfigFromJSON(document)
require.NoError(t, err, "a logged-out profile must still load")
require.Empty(t, reloaded.PrivateKey, "loading must not mint a key nothing will write down")
require.Empty(t, reloaded.SSHKey)
require.Equal(t, stored.ManagementURL.String(), reloaded.ManagementURL.String(),
"the rest of the profile survives the logout")
}
@@ -0,0 +1,131 @@
package profilemanager
import (
"encoding/json"
"os"
"path/filepath"
"reflect"
"testing"
"github.com/stretchr/testify/require"
)
// optionalBoolFields lists the *bool fields of Config by name, derived from the
// type so a field added later is covered without touching these tests.
func optionalBoolFields() []string {
pointerToBool := reflect.TypeOf((*bool)(nil))
var fields []string
configType := reflect.TypeOf(Config{})
for i := range configType.NumField() {
field := configType.Field(i)
if field.Type == pointerToBool && field.Tag.Get("json") != "-" {
fields = append(fields, field.Name)
}
}
return fields
}
func requireNoUnsetOptionalBool(t *testing.T, config *Config, context string) {
t.Helper()
value := reflect.ValueOf(*config)
for _, name := range optionalBoolFields() {
require.False(t, value.FieldByName(name).IsNil(),
"%s left %s unset, so its readers have to invent a default and a diff of it compares presence instead of value", context, name)
}
}
// An optional bool must not be tristate. While one can be nil, true or false,
// every reader has to invent the meaning of nil, and — the reason this test
// exists — a diff of the config ends up comparing presence rather than value:
// that is what made the update-settings gate refuse `netbird up` for a client
// restating its own defaults. apply() is where a config becomes complete, so
// the invariant belongs to it: no *bool may come out of apply() unset.
func TestApplyLeavesNoOptionalBoolUnset(t *testing.T) {
require.NotEmpty(t, optionalBoolFields(), "the invariant is only meaningful while Config has optional bools")
t.Run("a config built from scratch", func(t *testing.T) {
config := newConfigSkeleton()
_, err := config.apply(ConfigInput{})
require.NoError(t, err)
requireNoUnsetOptionalBool(t, config, "apply on a new config")
})
t.Run("a config file that predates every optional field", func(t *testing.T) {
path := filepath.Join(t.TempDir(), "legacy.json")
require.NoError(t, os.WriteFile(path, []byte(`{"WgIface":"wt0"}`), 0o600))
config, err := GetExistingConfig(path)
require.NoError(t, err)
requireNoUnsetOptionalBool(t, config, "a read of a legacy config")
})
t.Run("a config file that stores them as null", func(t *testing.T) {
path := filepath.Join(t.TempDir(), "null.json")
_, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path})
require.NoError(t, err)
unsetOnDisk(t, path, optionalBoolFields()...)
config, err := GetExistingConfig(path)
require.NoError(t, err)
requireNoUnsetOptionalBool(t, config, "a read of a config storing nulls")
})
}
// The same invariant on disk: what a write leaves in the file is what the next
// client to read it starts from, so no write may store a null.
func TestNoWriteStoresAnUnsetOptionalBool(t *testing.T) {
requireNoNullOnDisk := func(t *testing.T, path string, context string) {
t.Helper()
raw, err := os.ReadFile(path)
require.NoError(t, err)
var stored map[string]json.RawMessage
require.NoError(t, json.Unmarshal(raw, &stored))
for _, name := range optionalBoolFields() {
value, present := stored[name]
require.True(t, present, "%s did not store %s at all", context, name)
require.NotEqual(t, "null", string(value), "%s stored %s as null", context, name)
}
}
t.Run("UpdateOrCreateConfig", func(t *testing.T) {
path := filepath.Join(t.TempDir(), "created.json")
_, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path, ManagementURL: DefaultManagementURL})
require.NoError(t, err)
requireNoNullOnDisk(t, path, "UpdateOrCreateConfig")
})
t.Run("UpdateConfig over a config storing nulls", func(t *testing.T) {
path := filepath.Join(t.TempDir(), "stored.json")
_, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path})
require.NoError(t, err)
unsetOnDisk(t, path, optionalBoolFields()...)
_, err = UpdateConfig(ConfigInput{ConfigPath: path, ManagementURL: "https://mgmt.example.com"})
require.NoError(t, err)
requireNoNullOnDisk(t, path, "UpdateConfig")
})
// Renaming used to copy the file back through a bare Unmarshal, which
// preserved the nulls a pre-fix client had written.
t.Run("RenameProfile", func(t *testing.T) {
withTestSM(t, func(sm *ServiceManager, username string) {
created, err := sm.AddProfile("work", username)
require.NoError(t, err)
unsetOnDisk(t, created.Path, optionalBoolFields()...)
require.NoError(t, sm.RenameProfile(created.ID, username, "office"))
requireNoNullOnDisk(t, created.Path, "RenameProfile")
})
})
}
@@ -0,0 +1,96 @@
package profilemanager
import (
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/x509"
"crypto/x509/pkix"
"encoding/pem"
"math/big"
"os"
"path/filepath"
"testing"
"time"
"github.com/stretchr/testify/require"
)
// writeCertPair writes a throwaway certificate and key, so apply() has
// something real to load rather than a missing file it would only log about.
func writeCertPair(t *testing.T) (certPath, keyPath string) {
t.Helper()
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
require.NoError(t, err)
template := x509.Certificate{
SerialNumber: big.NewInt(1),
Subject: pkix.Name{CommonName: "probe-test"},
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(time.Hour),
}
der, err := x509.CreateCertificate(rand.Reader, &template, &template, &key.PublicKey, key)
require.NoError(t, err)
keyDER, err := x509.MarshalECPrivateKey(key)
require.NoError(t, err)
dir := t.TempDir()
certPath = filepath.Join(dir, "client.crt")
keyPath = filepath.Join(dir, "client.key")
require.NoError(t, os.WriteFile(certPath, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}), 0o600))
require.NoError(t, os.WriteFile(keyPath, pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}), 0o600))
return certPath, keyPath
}
// The dry run behind the update-settings gate must not read the mTLS pair off
// disk. The loaded pair feeds the connection, never the comparison, and the
// gate runs it on every SetConfig and Login — twice per request — including the
// ones it refuses.
func TestProbeDoesNotLoadTheCertificatePair(t *testing.T) {
certPath, keyPath := writeCertPair(t)
t.Run("a real apply loads it", func(t *testing.T) {
config := newConfigSkeleton()
config.ClientCertPath, config.ClientCertKeyPath = certPath, keyPath
_, err := config.apply(ConfigInput{})
require.NoError(t, err)
require.NotNil(t, config.ClientCertKeyPair, "the connection would have no client certificate")
})
t.Run("a probe does not", func(t *testing.T) {
config := newConfigSkeleton()
config.ClientCertPath, config.ClientCertKeyPath = certPath, keyPath
config.probing = true
_, err := config.apply(ConfigInput{})
require.NoError(t, err)
require.Nil(t, config.ClientCertKeyPair, "the dry run read the certificate off disk")
})
// And the verdict is the same either way, which is the only thing the gate
// asks of the probe.
t.Run("the verdict is unaffected", func(t *testing.T) {
path := filepath.Join(t.TempDir(), "mtls.json")
_, err := UpdateOrCreateConfig(ConfigInput{
ConfigPath: path,
ManagementURL: DefaultManagementURL,
ClientCertPath: certPath,
ClientCertKeyPath: keyPath,
})
require.NoError(t, err)
stored, err := GetExistingConfig(path)
require.NoError(t, err)
changed, err := stored.WouldChange(ConfigInput{ClientCertPath: certPath, ClientCertKeyPath: keyPath})
require.NoError(t, err)
require.False(t, changed, "restating the stored certificate paths is not a change")
changed, err = stored.WouldChange(ConfigInput{ClientCertPath: filepath.Join(t.TempDir(), "other.crt")})
require.NoError(t, err)
require.True(t, changed, "a different certificate path is a change")
})
}
@@ -196,7 +196,7 @@ func TestWireguardPortZeroExplicit(t *testing.T) {
assert.Equal(t, 0, config.WgPort, "WgPort should be 0 when explicitly set by user")
// Verify it persists
readConfig, err := GetConfig(configPath)
readConfig, err := GetExistingConfig(configPath)
require.NoError(t, err)
assert.Equal(t, 0, readConfig.WgPort, "WgPort should remain 0 after reading from file")
}
@@ -0,0 +1,529 @@
package profilemanager
import (
"encoding/json"
"os"
"path/filepath"
"runtime"
"testing"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/iface"
"github.com/netbirdio/netbird/shared/management/domain"
)
func seededConfig(t *testing.T) *Config {
t.Helper()
path := filepath.Join(t.TempDir(), "seeded.json")
cfg, err := UpdateOrCreateConfig(ConfigInput{
ConfigPath: path,
ManagementURL: "https://api.netbird.io:443",
PreSharedKey: strPointer("stored-key"),
})
require.NoError(t, err)
return cfg
}
func strPointer(s string) *string { return &s }
func intPtr(i int) *int { return &i }
func TestWouldChange(t *testing.T) {
tests := []struct {
name string
input ConfigInput
want bool
}{
{name: "empty input", input: ConfigInput{}, want: false},
{name: "same management URL", input: ConfigInput{ManagementURL: "https://api.netbird.io:443"}, want: false},
{name: "management URL without its default port", input: ConfigInput{ManagementURL: "https://api.netbird.io"}, want: false},
{name: "different management URL", input: ConfigInput{ManagementURL: "https://other.example:443"}, want: true},
{name: "same pre-shared key", input: ConfigInput{PreSharedKey: strPointer("stored-key")}, want: false},
{name: "redacted pre-shared key", input: ConfigInput{PreSharedKey: strPointer("**********")}, want: false},
{name: "different pre-shared key", input: ConfigInput{PreSharedKey: strPointer("other-key")}, want: true},
{name: "new interface blacklist entry", input: ConfigInput{ExtraIFaceBlackList: []string{"nb-probe0"}}, want: true},
{name: "blacklist entry already present", input: ConfigInput{ExtraIFaceBlackList: []string{"lo"}}, want: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := seededConfig(t)
changed, err := cfg.WouldChange(tt.input)
require.NoError(t, err)
require.Equal(t, tt.want, changed)
})
}
}
// The dry run must not be observable on the config it is run against: it
// decides whether a write is allowed, it does not perform one.
func TestWouldChangeLeavesTheConfigAlone(t *testing.T) {
cfg := seededConfig(t)
blacklist := len(cfg.IFaceBlackList)
changed, err := cfg.WouldChange(ConfigInput{
ManagementURL: "https://other.example:443",
PreSharedKey: strPointer("other-key"),
ExtraIFaceBlackList: []string{"nb-probe0"},
DNSLabels: domain.FromPunycodeList([]string{"probe"}),
NATExternalIPs: []string{"1.2.3.4"},
})
require.NoError(t, err)
require.True(t, changed)
require.Equal(t, "https://api.netbird.io:443", cfg.ManagementURL.String())
require.Equal(t, "stored-key", cfg.PreSharedKey)
require.Len(t, cfg.IFaceBlackList, blacklist)
require.Empty(t, cfg.DNSLabels)
require.Empty(t, cfg.NATExternalIPs)
}
// A nil config means the profile holds nothing yet, so the baseline is what
// the daemon would create for it.
func TestWouldChangeWithoutAStoredConfig(t *testing.T) {
var cfg *Config
changed, err := cfg.WouldChange(ConfigInput{})
require.NoError(t, err)
require.False(t, changed, "a request carrying nothing cannot change anything")
changed, err = cfg.WouldChange(ConfigInput{ManagementURL: DefaultManagementURL})
require.NoError(t, err)
require.False(t, changed, "the default management URL is what would be written anyway")
changed, err = cfg.WouldChange(ConfigInput{ManagementURL: "https://other.example:443"})
require.NoError(t, err)
require.True(t, changed)
}
func TestWouldChangeReportsAnInvalidInput(t *testing.T) {
cfg := seededConfig(t)
_, err := cfg.WouldChange(ConfigInput{ManagementURL: "not-a-url"})
require.Error(t, err)
}
// Reads must not write. A config file missing a field apply() fills in (MTU,
// here) is what used to trigger the write-back.
func TestReadsDoNotWriteTheConfigBack(t *testing.T) {
denormalized := []byte(`{"WgIface":"wt0"}`)
for name, read := range map[string]func(string) (*Config, error){
"GetExistingConfig": GetExistingConfig,
"ReadConfigOrDefault": ReadConfigOrDefault,
} {
t.Run(name, func(t *testing.T) {
path := filepath.Join(t.TempDir(), "profile.json")
require.NoError(t, os.WriteFile(path, denormalized, 0o600))
cfg, err := read(path)
require.NoError(t, err)
require.Equal(t, uint16(iface.DefaultMTU), cfg.MTU, "the returned config is still normalized in memory")
require.Empty(t, cfg.PrivateKey, "a read must not mint an identity either")
after, err := os.ReadFile(path)
require.NoError(t, err)
require.Equal(t, string(denormalized), string(after), "%s rewrote the config file", name)
})
}
}
// ReadConfigOrDefault resolves a default config for a profile that has no file
// yet, and that must not create the file either.
func TestReadConfigDoesNotCreateTheFile(t *testing.T) {
path := filepath.Join(t.TempDir(), "absent.json")
cfg, err := ReadConfigOrDefault(path)
require.NoError(t, err)
require.Equal(t, DefaultManagementURL, cfg.ManagementURL.String())
_, err = os.Stat(path)
require.True(t, os.IsNotExist(err), "ReadConfigOrDefault created the config file")
}
// The identity is the one thing a read cannot recompute, so it is provisioned
// on request and its caller persists it.
func TestEnsureIdentity(t *testing.T) {
cfg := newConfigSkeleton()
generated, err := cfg.EnsureIdentity()
require.NoError(t, err)
require.True(t, generated)
require.NotEmpty(t, cfg.PrivateKey)
require.NotEmpty(t, cfg.SSHKey)
key := cfg.PrivateKey
generated, err = cfg.EnsureIdentity()
require.NoError(t, err)
require.False(t, generated, "a config that already has an identity keeps it")
require.Equal(t, key, cfg.PrivateKey)
}
// One endpoint written several ways is one endpoint. A gate that compared
// spellings refused a client restating its own management URL with a trailing
// slash, which is a normal way to write it.
func TestSameServiceURL(t *testing.T) {
tests := []struct {
a, b string
want bool
}{
{a: "https://mgmt.example.com", b: "https://mgmt.example.com:443", want: true},
{a: "https://mgmt.example.com", b: "https://mgmt.example.com/", want: true},
{a: "https://mgmt.example.com/", b: "https://mgmt.example.com:443/", want: true},
{a: "https://MGMT.example.com", b: "https://mgmt.example.com", want: true},
{a: "http://mgmt.example.com", b: "http://mgmt.example.com:80", want: true},
{a: "https://mgmt.example.com", b: "http://mgmt.example.com", want: false},
{a: "https://mgmt.example.com", b: "https://mgmt.example.com:8443", want: false},
{a: "https://mgmt.example.com", b: "https://other.example.com", want: false},
}
for _, tt := range tests {
t.Run(tt.a+" vs "+tt.b, func(t *testing.T) {
a, err := ParseServiceURL("a", tt.a)
require.NoError(t, err)
b, err := ParseServiceURL("b", tt.b)
require.NoError(t, err)
require.Equal(t, tt.want, SameServiceURL(a, b))
require.Equal(t, tt.want, SameServiceURL(b, a), "the comparison must be symmetric")
})
}
}
// The same spellings, through the dry run the update-settings gate uses.
func TestWouldChangeIgnoresURLSpelling(t *testing.T) {
path := filepath.Join(t.TempDir(), "seeded.json")
_, err := UpdateOrCreateConfig(ConfigInput{
ConfigPath: path,
ManagementURL: "https://mgmt.example.com",
})
require.NoError(t, err)
cfg, err := GetExistingConfig(path)
require.NoError(t, err)
for _, spelling := range []string{
"https://mgmt.example.com",
"https://mgmt.example.com/",
"https://mgmt.example.com:443",
"https://mgmt.example.com:443/",
"https://MGMT.example.com",
} {
changed, err := cfg.WouldChange(ConfigInput{ManagementURL: spelling})
require.NoError(t, err)
require.False(t, changed, "%q is the stored endpoint written differently", spelling)
}
changed, err := cfg.WouldChange(ConfigInput{ManagementURL: "https://mgmt.example.com:8443"})
require.NoError(t, err)
require.True(t, changed, "a different port is a different endpoint")
}
// The dry-run baseline exists to be compared against and discarded, so it must
// not mint keys — the CLI's login backoff loop would otherwise log a fresh
// "generated new Wireguard key" on every attempt.
func TestDryRunBaselineDoesNotGenerateKeys(t *testing.T) {
baseline, err := newDryRunBaseline(filepath.Join(t.TempDir(), "absent.json"))
require.NoError(t, err)
require.Empty(t, baseline.PrivateKey, "generated a WireGuard key for a throwaway config")
require.Empty(t, baseline.SSHKey, "generated an SSH key for a throwaway config")
// Everything the comparison actually looks at is still the default config.
require.Equal(t, DefaultManagementURL, baseline.ManagementURL.String())
require.Equal(t, uint16(iface.DefaultMTU), baseline.MTU)
require.Equal(t, iface.DefaultWgPort, baseline.WgPort)
}
// A stored profile can carry no identity — a mobile logout clears the keys in
// place — so the next config write has to mint one, which is what keeps the
// following login from dialing management with an empty key.
func TestUpdateConfigProvisionsAMissingIdentity(t *testing.T) {
path := filepath.Join(t.TempDir(), "logged-out.json")
_, err := UpdateOrCreateConfig(ConfigInput{
ConfigPath: path,
ManagementURL: "https://api.netbird.io:443",
})
require.NoError(t, err)
// Stand in for the logout, which zeroes the keys and writes the config out.
loggedOut, err := GetExistingConfig(path)
require.NoError(t, err)
loggedOut.PrivateKey = ""
loggedOut.SSHKey = ""
require.NoError(t, WriteOutConfig(path, loggedOut))
cfg, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path})
require.NoError(t, err)
require.NotEmpty(t, cfg.PrivateKey, "the write path did not provision an identity")
require.NotEmpty(t, cfg.SSHKey)
persisted, err := GetExistingConfig(path)
require.NoError(t, err)
require.Equal(t, cfg.PrivateKey, persisted.PrivateKey, "the provisioned identity was not persisted")
}
// A config that carries no sync message version must not make the dry run
// panic: the gate runs inside a request handler, where failing closed is the
// worst acceptable outcome.
func TestWouldChangeWithoutAStoredSyncMessageVersion(t *testing.T) {
cfg := seededConfig(t)
require.Nil(t, cfg.SyncMessageVersion, "the fixture is only useful while the field starts out unset")
version := 2
changed, err := cfg.WouldChange(ConfigInput{SyncMessageVersion: &version})
require.NoError(t, err)
require.True(t, changed)
require.Nil(t, cfg.SyncMessageVersion, "the dry run set the version on the stored config")
}
// Restating the certificate paths a config already holds is not a change, for
// the same reason restating any other value is not.
func TestWouldChangeIgnoresRestatedCertificatePaths(t *testing.T) {
path := filepath.Join(t.TempDir(), "mtls.json")
_, err := UpdateOrCreateConfig(ConfigInput{
ConfigPath: path,
ManagementURL: "https://api.netbird.io:443",
ClientCertPath: "/etc/netbird/client.crt",
ClientCertKeyPath: "/etc/netbird/client.key",
})
require.NoError(t, err)
cfg, err := GetExistingConfig(path)
require.NoError(t, err)
changed, err := cfg.WouldChange(ConfigInput{
ClientCertPath: "/etc/netbird/client.crt",
ClientCertKeyPath: "/etc/netbird/client.key",
})
require.NoError(t, err)
require.False(t, changed, "the stored certificate paths were restated")
changed, err = cfg.WouldChange(ConfigInput{ClientCertPath: "/etc/netbird/other.crt"})
require.NoError(t, err)
require.True(t, changed, "a different certificate path is a change")
}
// A read that lands on a missing file must not hand back keys: nothing would
// write them down, so the caller would connect with an identity that changes on
// the next run and registers a second peer.
func TestReadConfigOrDefaultCarriesNoIdentity(t *testing.T) {
cfg, err := ReadConfigOrDefault(filepath.Join(t.TempDir(), "absent.json"))
require.NoError(t, err)
require.Empty(t, cfg.PrivateKey, "a read minted a WireGuard key")
require.Empty(t, cfg.SSHKey, "a read minted an SSH key")
// So the caller's own EnsureIdentity is the one that reports the work, and
// therefore the one that triggers the write.
generated, err := cfg.EnsureIdentity()
require.NoError(t, err)
require.True(t, generated, "the provisioning caller could not tell it had to persist the identity")
}
// CreateInMemoryConfig is the opposite contract: its callers connect with what
// they get back, so it does carry an identity.
func TestCreateInMemoryConfigCarriesAnIdentity(t *testing.T) {
cfg, err := CreateInMemoryConfig(ConfigInput{ManagementURL: "https://api.netbird.io:443"})
require.NoError(t, err)
require.NotEmpty(t, cfg.PrivateKey)
require.NotEmpty(t, cfg.SSHKey)
}
// The admin panel is opened, not dialed, so its path identifies it. Comparing
// it as a bare endpoint left a custom panel URL unable to change.
func TestAdminURLPathIsPartOfTheIdentity(t *testing.T) {
path := filepath.Join(t.TempDir(), "panel.json")
_, err := UpdateOrCreateConfig(ConfigInput{
ConfigPath: path,
AdminURL: "https://app.example.com/netbird",
})
require.NoError(t, err)
cfg, err := GetExistingConfig(path)
require.NoError(t, err)
require.Equal(t, "https://app.example.com:443/netbird", cfg.AdminURL.String())
// Equivalent spellings of the same panel are still not a change.
for _, same := range []string{
"https://app.example.com/netbird",
"https://app.example.com:443/netbird",
"https://app.example.com/netbird/",
"https://APP.example.com/netbird",
} {
changed, err := cfg.WouldChange(ConfigInput{AdminURL: same})
require.NoError(t, err)
require.False(t, changed, "%q is the stored panel written differently", same)
}
// A different path is a different panel, and it must be persisted.
changed, err := cfg.WouldChange(ConfigInput{AdminURL: "https://app.example.com/other"})
require.NoError(t, err)
require.True(t, changed, "a different panel path is a change")
updated, err := UpdateConfig(ConfigInput{ConfigPath: path, AdminURL: "https://app.example.com/other"})
require.NoError(t, err)
require.Equal(t, "https://app.example.com:443/other", updated.AdminURL.String(), "the new panel path was not persisted")
}
// unsetOnDisk rewrites the stored config so the named fields carry a JSON null,
// which is how a profile written before apply() resolved them looks on disk.
// It synthesizes that state: no write produces it any more.
func unsetOnDisk(t *testing.T, path string, fields ...string) {
t.Helper()
raw, err := os.ReadFile(path)
require.NoError(t, err)
var stored map[string]json.RawMessage
require.NoError(t, json.Unmarshal(raw, &stored))
for _, field := range fields {
_, present := stored[field]
require.True(t, present, "%s is not a field of the stored config", field)
stored[field] = json.RawMessage("null")
}
rewritten, err := json.Marshal(stored)
require.NoError(t, err)
require.NoError(t, os.WriteFile(path, rewritten, 0600))
}
// Seven fields mean "the effective default" when they hold no value, and every
// profile written before apply() resolved them holds them as null. Restating
// that default is asking for no change — and the CLI restates it on every
// `netbird up`, because a flag set through an environment variable is a flag
// pflag reports as Changed. Judging those restatements as changes made the
// update-settings gate refuse `netbird up` outright for a client configured
// through the environment, which is the shape of a Kubernetes deployment.
//
// A login now writes those fields set, so the fixture puts the null state back
// on disk with unsetOnDisk instead of getting it from a login.
func TestWouldChangeIgnoresRestatedDefaultsOfUnsetFields(t *testing.T) {
networkMonitorDefault := runtime.GOOS == "windows" || runtime.GOOS == "darwin"
tests := []struct {
field string
theDefault ConfigInput
theOtherWay ConfigInput
}{
{"EnableSSHRoot",
ConfigInput{EnableSSHRoot: boolPtr(false)}, ConfigInput{EnableSSHRoot: boolPtr(true)}},
{"EnableSSHSFTP",
ConfigInput{EnableSSHSFTP: boolPtr(false)}, ConfigInput{EnableSSHSFTP: boolPtr(true)}},
{"EnableSSHLocalPortForwarding",
ConfigInput{EnableSSHLocalPortForwarding: boolPtr(false)}, ConfigInput{EnableSSHLocalPortForwarding: boolPtr(true)}},
{"EnableSSHRemotePortForwarding",
ConfigInput{EnableSSHRemotePortForwarding: boolPtr(false)}, ConfigInput{EnableSSHRemotePortForwarding: boolPtr(true)}},
{"DisableSSHAuth",
ConfigInput{DisableSSHAuth: boolPtr(false)}, ConfigInput{DisableSSHAuth: boolPtr(true)}},
{"SSHJWTCacheTTL",
ConfigInput{SSHJWTCacheTTL: intPtr(0)}, ConfigInput{SSHJWTCacheTTL: intPtr(300)}},
{"NetworkMonitor",
ConfigInput{NetworkMonitor: boolPtr(networkMonitorDefault)}, ConfigInput{NetworkMonitor: boolPtr(!networkMonitorDefault)}},
}
for _, tt := range tests {
t.Run(tt.field, func(t *testing.T) {
path := filepath.Join(t.TempDir(), "unset.json")
_, err := UpdateOrCreateConfig(ConfigInput{
ConfigPath: path,
ManagementURL: "https://api.netbird.io:443",
})
require.NoError(t, err)
unsetOnDisk(t, path, tt.field)
cfg, err := GetExistingConfig(path)
require.NoError(t, err)
changed, err := cfg.WouldChange(tt.theDefault)
require.NoError(t, err)
require.False(t, changed, "restating the default of an unset %s was judged a change", tt.field)
// The gate still has to refuse a request that does ask for something.
changed, err = cfg.WouldChange(tt.theOtherWay)
require.NoError(t, err)
require.True(t, changed, "asking for a non-default %s is a change", tt.field)
})
}
}
// The verdict must not depend on where the caller got the config from. Readers
// normalize what they hand out, but apply() signals "I filled in a default"
// through the same bool as "the input changed something", so a config that
// never passed through a read would otherwise report a change for an input
// that asks for nothing.
func TestWouldChangeNormalizesBeforeMeasuring(t *testing.T) {
rawConfig := func(t *testing.T) *Config {
t.Helper()
cfg := &Config{WgIface: iface.WgInterfaceDefault}
require.Nil(t, cfg.ServerSSHAllowed, "the fixture is only useful while the config is not normalized")
require.Nil(t, cfg.EnableSSHRoot)
require.Empty(t, cfg.IFaceBlackList)
return cfg
}
changed, err := rawConfig(t).WouldChange(ConfigInput{})
require.NoError(t, err)
require.False(t, changed, "an input carrying nothing cannot change anything")
changed, err = rawConfig(t).WouldChange(ConfigInput{EnableSSHRoot: boolPtr(false)})
require.NoError(t, err)
require.False(t, changed, "the default of a field the config never held is not a change")
changed, err = rawConfig(t).WouldChange(ConfigInput{EnableSSHRoot: boolPtr(true)})
require.NoError(t, err)
require.True(t, changed, "a non-default value is still a change")
}
// A zero-padded port addresses the same port. The normalization itself belongs
// to util.ServiceURLPort and is tested there; this asserts that the comparison
// this package hands its callers inherits it.
func TestServiceURLPortIsNormalizedNumerically(t *testing.T) {
padded, err := ParseServiceURL("padded", "https://mgmt.example.com:0443")
require.NoError(t, err)
plain, err := ParseServiceURL("plain", "https://mgmt.example.com:443")
require.NoError(t, err)
require.True(t, SameServiceURL(padded, plain))
}
// A list the profile does not have and a list the request empties are the same
// thing: no NAT mappings, no DNS labels. The profile stores an absent list as
// JSON null and reads it back as a nil slice, while `netbird up` sends the
// emptied list — CleanNATExternalIPs / CleanDNSLabels — whenever the matching
// environment variable is set to nothing, which a deployment template does by
// default. Judging nil and empty as different made the gate refuse that start,
// which is the very deadlock this branch exists to remove, on another field.
func TestWouldChangeIgnoresAnEmptiedListThatWasAlreadyAbsent(t *testing.T) {
path := filepath.Join(t.TempDir(), "lists.json")
_, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path, ManagementURL: DefaultManagementURL})
require.NoError(t, err)
stored, err := GetExistingConfig(path)
require.NoError(t, err)
require.Nil(t, stored.NATExternalIPs, "the fixture is only useful while the stored list is absent")
require.Nil(t, stored.DNSLabels)
changed, err := stored.WouldChange(ConfigInput{NATExternalIPs: make([]string, 0)})
require.NoError(t, err)
require.False(t, changed, "emptying a NAT list the profile never had is not a change")
changed, err = stored.WouldChange(ConfigInput{DNSLabels: domain.List{}})
require.NoError(t, err)
require.False(t, changed, "emptying a DNS label list the profile never had is not a change")
// A list that does hold something still moves when the request empties it.
withEntries, err := UpdateConfig(ConfigInput{ConfigPath: path, NATExternalIPs: []string{"1.2.3.4"}})
require.NoError(t, err)
require.Equal(t, []string{"1.2.3.4"}, withEntries.NATExternalIPs)
changed, err = withEntries.WouldChange(ConfigInput{NATExternalIPs: make([]string, 0)})
require.NoError(t, err)
require.True(t, changed, "clearing a NAT list that had an entry is a change")
}
+25 -8
View File
@@ -313,7 +313,11 @@ func (s *ServiceManager) AddProfile(displayName, username string) (*Profile, err
}
profPath := filepath.Join(configDir, id.String()+".json")
cfg, err := createNewConfig(ConfigInput{ConfigPath: profPath})
// Provisioned, not bare: this config goes straight to disk, and a profile
// file with no identity is one whose first reader has to mint the keys and
// remember to write them back. Before identity generation moved out of
// apply() into EnsureIdentity, createNewConfig produced them here too.
cfg, err := createProvisionedConfig(ConfigInput{ConfigPath: profPath})
if err != nil {
return nil, fmt.Errorf("failed to create new config: %w", err)
}
@@ -330,6 +334,19 @@ func (s *ServiceManager) AddProfile(displayName, username string) (*Profile, err
}, nil
}
// RenameProfile changes a profile's display name. It rewrites the whole
// profile file, not just the name: the config is read through the normalizing
// reader, so apply()'s resolved values — the optional booleans, the interface
// blacklist, the DNS route interval — are persisted along with the new name.
//
// That is deliberate. A write that skipped apply() is what left profiles on
// disk carrying null where a value was meant, and made a diff of the config
// compare presence instead of value. Two consequences worth knowing: the
// platform-dependent defaults resolved here are the renaming host's
// (ServerSSHAllowed and the network monitor differ per OS), and a profile
// whose stored name does not survive sanitizeDisplayName now fails to rename
// rather than being rewritten — though apply() rejects such a profile on every
// other read too, so it was already unusable.
func (s *ServiceManager) RenameProfile(id ID, username string, newName string) error {
displayName, err := sanitizeDisplayName(newName)
if err != nil {
@@ -356,17 +373,17 @@ func (s *ServiceManager) RenameProfile(id ID, username string, newName string) e
return ErrProfileNotFound
}
data, err := os.ReadFile(target.Path)
// Through the reader, not a bare Unmarshal: this was the one write that
// skipped apply(), so it copied back whatever the file held — including an
// optional field left unset, which every other write resolves to its
// default. Renaming a profile is a poor place to leave that behind.
cfg, err := GetExistingConfig(target.Path)
if err != nil {
return err
}
var cfg Config
if err := json.Unmarshal(data, &cfg); err != nil {
return err
return fmt.Errorf("read profile config: %w", err)
}
cfg.Name = displayName
if err := util.WriteJson(context.Background(), target.Path, cfg); err != nil {
if err := WriteOutConfig(target.Path, cfg); err != nil {
return fmt.Errorf("failed to write profile name: %w", err)
}
return nil
@@ -228,3 +228,27 @@ func TestRemoveProfile_DeletesStateFile(t *testing.T) {
assert.True(t, errors.Is(err, os.ErrNotExist), "state file should be removed")
})
}
// A profile file is written here and read back by whoever connects with it, so
// it has to carry the peer's identity. While AddProfile used the bare
// constructor, it wrote a config with no keys: the first reader had to mint
// them, and the paths that read without writing — a gate deciding whether to
// refuse a request, the mobile SDKs loading a stored profile — got a config
// that cannot connect.
func TestAddProfileWritesAnIdentity(t *testing.T) {
withTestSM(t, func(sm *ServiceManager, username string) {
created, err := sm.AddProfile("work", username)
require.NoError(t, err)
stored, err := GetExistingConfig(created.Path)
require.NoError(t, err)
require.NotEmpty(t, stored.PrivateKey, "the profile was written without a WireGuard key")
require.NotEmpty(t, stored.SSHKey, "the profile was written without an SSH key")
// And the identity is the one on disk, not one minted per read.
reread, err := GetExistingConfig(created.Path)
require.NoError(t, err)
require.Equal(t, stored.PrivateKey, reread.PrivateKey)
})
}
@@ -19,8 +19,7 @@ type IPForwardingState struct {
// routingV4/routingV6 track whether the routing path currently holds a
// reference, so repeated EnableRouting calls (one per network-map update)
// hold at most one reference per family and an unpaired DisableRouting
// can't release references held by DNAT rules.
// hold at most one reference per family.
routingV4 bool
routingV6 bool
@@ -95,31 +94,6 @@ func (f *IPForwardingState) ReleaseRouting() error {
return nil
}
// RequestForwarding enables the family's forwarding sysctl on first request.
func (f *IPForwardingState) RequestForwarding(v6 bool) error {
f.mu.Lock()
defer f.mu.Unlock()
if v6 {
return f.requestV6()
}
return f.requestV4()
}
// ReleaseForwarding decrements the family counter. The last v6 release restores
// what enable captured. v4 stays on: net.ipv4.ip_forward is co-owned by other
// tooling (docker, k8s, libvirt).
func (f *IPForwardingState) ReleaseForwarding(v6 bool) error {
f.mu.Lock()
defer f.mu.Unlock()
if v6 {
return f.releaseV6()
}
f.releaseV4()
return nil
}
func (f *IPForwardingState) requestV4() error {
if f.v4Count == 0 {
if err := systemops.EnableV4IPForwarding(); err != nil {
@@ -10,8 +10,7 @@ import (
)
// TestRequestRoutingV6ToV4Transition verifies that a v4-only routing request
// releases a previously held routing-owned v6 reference without touching
// references held by DNAT rules.
// releases a previously held routing-owned v6 reference.
func TestRequestRoutingV6ToV4Transition(t *testing.T) {
f := NewIPForwardingState("wt-fwd-test")
@@ -25,13 +24,6 @@ func TestRequestRoutingV6ToV4Transition(t *testing.T) {
assert.Equal(t, 1, v4, "v4 reference kept")
assert.Equal(t, 0, v6, "routing-owned v6 reference released")
// A DNAT-held reference survives a v4-only routing request.
require.NoError(t, f.RequestForwarding(true), "dnat v6 reference")
require.NoError(t, f.RequestRouting(false), "repeat v4-only request")
_, v6 = f.Counts()
assert.Equal(t, 1, v6, "dnat-held v6 reference survives")
require.NoError(t, f.ReleaseForwarding(true), "release dnat v6 reference")
require.NoError(t, f.ReleaseRouting(), "release routing")
v4, v6 = f.Counts()
assert.Equal(t, 0, v4, "all v4 references released")