mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-08 14:39:09 +02:00
Merge remote-tracking branch 'origin/main' into fix/pkce-flow-session-extend
# Conflicts: # shared/management/proto/management.pb.go
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
|
||||
@@ -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).
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user