diff --git a/client/internal/connect.go b/client/internal/connect.go index 88d829d2f..b42ef9818 100644 --- a/client/internal/connect.go +++ b/client/internal/connect.go @@ -292,6 +292,10 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan return nil } + if c.updateManager != nil { + c.updateManager.ResetMode() + } + // suspend connection attempts while the OS reports no usable network if waited, err := c.netMgr.Wait(c.ctx); err != nil { return nil diff --git a/client/internal/engine.go b/client/internal/engine.go index 4d731cbd7..e167f7589 100644 --- a/client/internal/engine.go +++ b/client/internal/engine.go @@ -455,7 +455,7 @@ func (e *Engine) stopLocked() { } if e.updateManager != nil { - e.updateManager.SetDownloadOnly() + e.updateManager.ResetMode() } log.Info("cleaning up status recorder states") @@ -951,11 +951,13 @@ func (e *Engine) handleAutoUpdateVersion(autoUpdateSettings *mgmProto.AutoUpdate } if autoUpdateSettings == nil { + log.Infof("no auto-update settings received, defaulting to download-only") + e.updateManager.SetDownloadOnly() return } if autoUpdateSettings.Version == disableAutoUpdate { - log.Infof("auto-update is disabled") + log.Infof("auto-update is disabled, switching to download-only") e.updateManager.SetDownloadOnly() return } diff --git a/client/internal/updater/manager.go b/client/internal/updater/manager.go index 1b69368d0..c730713d3 100644 --- a/client/internal/updater/manager.go +++ b/client/internal/updater/manager.go @@ -21,8 +21,16 @@ const ( latestVersion = "latest" ) +const ( + modeUndecided updateMode = iota + modeDownloadOnly + modeManaged +) + var errNoUpdateState = errors.New("no update state found") +type updateMode int + type UpdateState struct { PreUpdateVersion string TargetVersion string @@ -36,8 +44,9 @@ type Manager struct { statusRecorder *peer.Status stateManager *statemanager.Manager - downloadOnly bool // true when no enforcement from management; notifies UI to download latest - forceUpdate bool // true when management sets AlwaysUpdate; skips UI interaction and installs directly + mode updateMode + modeGen uint64 + forceUpdate bool // true when management sets AlwaysUpdate; skips UI interaction and installs directly lastTrigger time.Time mgmUpdateChan chan struct{} @@ -54,7 +63,7 @@ type Manager struct { pendingVersion *v.Version // updateMutex protects update, expectedVersion, updateToLatestVersion, - // downloadOnly, forceUpdate, pendingVersion, and lastTrigger fields + // mode, modeGen, forceUpdate, pendingVersion, and lastTrigger fields updateMutex sync.Mutex // installMutex and installing guard against concurrent installation attempts @@ -76,7 +85,6 @@ func NewManager(statusRecorder *peer.Status, stateManager *statemanager.Manager) updateChannel: make(chan struct{}, 1), currentVersion: version.NetbirdVersion(), update: version.NewUpdate("nb/client"), - downloadOnly: true, autoUpdateSupported: isAutoUpdateSupported, } @@ -151,11 +159,7 @@ func (m *Manager) Start(ctx context.Context) { func (m *Manager) SetDownloadOnly() { m.updateMutex.Lock() - m.downloadOnly = true - m.forceUpdate = false - m.expectedVersion = nil - m.updateToLatestVersion = false - m.lastTrigger = time.Time{} + m.setModeLocked(modeDownloadOnly) m.updateMutex.Unlock() select { @@ -169,6 +173,7 @@ func (m *Manager) SetVersion(expectedVersion string, forceUpdate bool) { if !m.autoUpdateSupported() { log.Warnf("auto-update not supported on this platform") + m.SetDownloadOnly() return } @@ -177,30 +182,32 @@ func (m *Manager) SetVersion(expectedVersion string, forceUpdate bool) { if expectedVersion == "" { log.Errorf("empty expected version provided") - m.expectedVersion = nil - m.updateToLatestVersion = false - m.downloadOnly = true + m.setModeLocked(modeDownloadOnly) return } - if expectedVersion == latestVersion { - m.updateToLatestVersion = true - m.expectedVersion = nil - } else { - expectedSemVer, err := v.NewVersion(expectedVersion) + var expectedSemVer *v.Version + if expectedVersion != latestVersion { + parsed, err := v.NewVersion(expectedVersion) if err != nil { - log.Errorf("error parsing version: %v", err) + log.Errorf("error parsing version, switching to download-only: %v", err) + m.setModeLocked(modeDownloadOnly) + select { + case m.mgmUpdateChan <- struct{}{}: + default: + } return } - if m.expectedVersion != nil && m.expectedVersion.Equal(expectedSemVer) { - return - } - m.expectedVersion = expectedSemVer - m.updateToLatestVersion = false + expectedSemVer = parsed } - m.lastTrigger = time.Time{} - m.downloadOnly = false + if m.sameDirectiveLocked(expectedSemVer, forceUpdate) { + return + } + + m.setModeLocked(modeManaged) + m.expectedVersion = expectedSemVer + m.updateToLatestVersion = expectedSemVer == nil m.forceUpdate = forceUpdate select { @@ -209,6 +216,13 @@ func (m *Manager) SetVersion(expectedVersion string, forceUpdate bool) { } } +func (m *Manager) ResetMode() { + m.updateMutex.Lock() + defer m.updateMutex.Unlock() + + m.setModeLocked(modeUndecided) +} + // Install triggers the installation of the pending version. It is called when the user clicks the install button in the UI. func (m *Manager) Install(ctx context.Context) error { if !m.autoUpdateSupported() { @@ -255,12 +269,17 @@ func (m *Manager) NotifyUI() { m.updateMutex.Unlock() return } - downloadOnly := m.downloadOnly + mode := m.mode + gen := m.modeGen pendingVersion := m.pendingVersion latestVersion := m.update.LatestVersion() m.updateMutex.Unlock() - if downloadOnly { + if mode == modeUndecided { + return + } + + if mode == modeDownloadOnly { if latestVersion == nil { return } @@ -268,6 +287,9 @@ func (m *Manager) NotifyUI() { if err != nil || currentVersion.GreaterThanOrEqual(latestVersion) { return } + if m.modeChanged(gen) { + return + } m.statusRecorder.PublishEvent( cProto.SystemEvent_INFO, cProto.SystemEvent_SYSTEM, @@ -278,7 +300,7 @@ func (m *Manager) NotifyUI() { return } - if pendingVersion != nil { + if pendingVersion != nil && !m.modeChanged(gen) { m.statusRecorder.PublishEvent( cProto.SystemEvent_INFO, cProto.SystemEvent_SYSTEM, @@ -343,13 +365,18 @@ func (m *Manager) handleUpdate(ctx context.Context) { return } - downloadOnly := m.downloadOnly + mode := m.mode + gen := m.modeGen forceUpdate := m.forceUpdate curLatestVersion := m.update.LatestVersion() switch { + case mode == modeUndecided: + log.Tracef("auto-update mode not decided yet") + m.updateMutex.Unlock() + return // Download-only mode or resolve "latest" to actual version - case downloadOnly, m.updateToLatestVersion: + case mode == modeDownloadOnly, m.updateToLatestVersion: if curLatestVersion == nil { log.Tracef("latest version not fetched yet") m.updateMutex.Unlock() @@ -374,12 +401,17 @@ func (m *Manager) handleUpdate(ctx context.Context) { m.lastTrigger = time.Now() log.Infof("new version available: %s", updateVersion) - if !downloadOnly && !forceUpdate { + if mode == modeManaged && !forceUpdate { m.pendingVersion = updateVersion } m.updateMutex.Unlock() - if downloadOnly { + if m.modeChanged(gen) { + log.Debugf("auto-update mode changed while checking %s, discarding", updateVersion) + return + } + + if mode == modeDownloadOnly { m.statusRecorder.PublishEvent( cProto.SystemEvent_INFO, cProto.SystemEvent_SYSTEM, @@ -406,6 +438,33 @@ func (m *Manager) handleUpdate(ctx context.Context) { ) } +func (m *Manager) modeChanged(gen uint64) bool { + m.updateMutex.Lock() + defer m.updateMutex.Unlock() + + return m.modeGen != gen +} + +func (m *Manager) sameDirectiveLocked(expectedVersion *v.Version, forceUpdate bool) bool { + if m.mode != modeManaged || m.forceUpdate != forceUpdate { + return false + } + if expectedVersion == nil { + return m.updateToLatestVersion + } + return m.expectedVersion != nil && m.expectedVersion.Equal(expectedVersion) +} + +func (m *Manager) setModeLocked(mode updateMode) { + m.mode = mode + m.modeGen++ + m.forceUpdate = false + m.expectedVersion = nil + m.updateToLatestVersion = false + m.pendingVersion = nil + m.lastTrigger = time.Time{} +} + func (m *Manager) install(ctx context.Context, pendingVersion *v.Version) error { m.statusRecorder.PublishEvent( cProto.SystemEvent_CRITICAL, diff --git a/client/internal/updater/manager_linux_test.go b/client/internal/updater/manager_linux_test.go index b05dd7e7d..4501ddde7 100644 --- a/client/internal/updater/manager_linux_test.go +++ b/client/internal/updater/manager_linux_test.go @@ -16,7 +16,7 @@ import ( ) // On Linux, only Mode 1 (downloadOnly) is supported. -// SetVersion is a no-op because auto-update installation is not supported. +// SetVersion falls back to download-only because auto-update installation is not supported. func Test_LatestVersion_Linux(t *testing.T) { testMatrix := []struct { @@ -70,7 +70,7 @@ func Test_LatestVersion_Linux(t *testing.T) { t.Errorf("%s: Initial version mismatch, expected %v, got %v", c.name, c.initialLatestVersion.String(), ver) } - mockUpdate.latestVersion = c.latestVersion + mockUpdate.setLatestVersion(c.latestVersion) mockUpdate.onUpdate() ver, enforced = waitForUpdateEvent(sub, 500*time.Millisecond) @@ -89,22 +89,24 @@ func Test_LatestVersion_Linux(t *testing.T) { } } -func Test_SetVersion_NoOp_Linux(t *testing.T) { - // On Linux, SetVersion should be a no-op — no events fired - tmpFile := path.Join(t.TempDir(), "update-test-noop.json") +func Test_SetVersion_FallsBackToDownloadOnly_Linux(t *testing.T) { + tmpFile := path.Join(t.TempDir(), "update-test-fallback.json") recorder := peer.NewRecorder("") sub := recorder.SubscribeToEvents() defer recorder.UnsubscribeFromEvents(sub) m := NewManager(recorder, statemanager.New(tmpFile)) - m.update = &versionUpdateMock{latestVersion: v.Must(v.NewSemver("1.0.1"))} + m.update = &versionUpdateMock{latestVersion: v.Must(v.NewSemver("1.0.5"))} m.currentVersion = "1.0.0" m.Start(context.Background()) m.SetVersion("1.0.1", false) - ver, _ := waitForUpdateEvent(sub, 500*time.Millisecond) - if ver != "" { - t.Errorf("SetVersion should be a no-op on Linux, but got event with version %s", ver) + ver, enforced := waitForUpdateEvent(sub, 500*time.Millisecond) + if ver != "1.0.5" { + t.Fatalf("expected download-only event for fetched 1.0.5, got %q", ver) + } + if enforced { + t.Error("Linux fallback must never have enforced metadata") } m.Stop() diff --git a/client/internal/updater/manager_mode_test.go b/client/internal/updater/manager_mode_test.go new file mode 100644 index 000000000..aa238f94c --- /dev/null +++ b/client/internal/updater/manager_mode_test.go @@ -0,0 +1,217 @@ +package updater + +import ( + "context" + "path" + "testing" + "time" + + v "github.com/hashicorp/go-version" + + "github.com/netbirdio/netbird/client/internal/peer" + "github.com/netbirdio/netbird/client/internal/statemanager" +) + +func Test_UndecidedMode_SuppressesNotification(t *testing.T) { + tmpFile := path.Join(t.TempDir(), "update-test-undecided.json") + recorder := peer.NewRecorder("") + sub := recorder.SubscribeToEvents() + defer recorder.UnsubscribeFromEvents(sub) + + mockUpdate := &versionUpdateMock{latestVersion: v.Must(v.NewSemver("1.0.1"))} + m := NewManager(recorder, statemanager.New(tmpFile)) + m.update = mockUpdate + m.currentVersion = "1.0.0" + m.Start(context.Background()) + defer m.Stop() + + mockUpdate.onUpdate() + if ver, _ := waitForUpdateEvent(sub, 300*time.Millisecond); ver != "" { + t.Fatalf("undecided mode must not publish, got %q", ver) + } + + m.NotifyUI() + if ver, _ := waitForUpdateEvent(sub, 300*time.Millisecond); ver != "" { + t.Fatalf("NotifyUI in undecided mode must not publish, got %q", ver) + } + + m.SetDownloadOnly() + ver, enforced := waitForUpdateEvent(sub, 500*time.Millisecond) + if ver != "1.0.1" { + t.Fatalf("expected download-only event for 1.0.1, got %q", ver) + } + if enforced { + t.Error("download-only event must not carry enforced metadata") + } +} + +func Test_ResetMode_ReturnsToUndecided(t *testing.T) { + tmpFile := path.Join(t.TempDir(), "update-test-reset.json") + recorder := peer.NewRecorder("") + sub := recorder.SubscribeToEvents() + defer recorder.UnsubscribeFromEvents(sub) + + mockUpdate := &versionUpdateMock{latestVersion: v.Must(v.NewSemver("1.0.1"))} + m := NewManager(recorder, statemanager.New(tmpFile)) + m.update = mockUpdate + m.currentVersion = "1.0.0" + m.autoUpdateSupported = func() bool { return true } + m.Start(context.Background()) + defer m.Stop() + + m.SetVersion("1.0.1", false) + ver, enforced := waitForUpdateEvent(sub, 500*time.Millisecond) + if ver != "1.0.1" || !enforced { + t.Fatalf("expected enforced event for 1.0.1, got %q enforced=%v", ver, enforced) + } + + m.ResetMode() + + mockUpdate.onUpdate() + if ver, _ := waitForUpdateEvent(sub, 300*time.Millisecond); ver != "" { + t.Fatalf("reset mode must not publish on fetch, got %q", ver) + } + + m.NotifyUI() + if ver, _ := waitForUpdateEvent(sub, 300*time.Millisecond); ver != "" { + t.Fatalf("NotifyUI after reset must not publish, got %q", ver) + } + + if err := m.Install(context.Background()); err == nil { + t.Fatal("Install after reset must fail without a pending version") + } + + m.SetVersion("1.0.1", false) + ver, enforced = waitForUpdateEvent(sub, 500*time.Millisecond) + if ver != "1.0.1" || !enforced { + t.Fatalf("expected enforced event again after reset, got %q enforced=%v", ver, enforced) + } +} + +func Test_SetDownloadOnly_ClearsPendingVersion(t *testing.T) { + tmpFile := path.Join(t.TempDir(), "update-test-pending.json") + recorder := peer.NewRecorder("") + sub := recorder.SubscribeToEvents() + defer recorder.UnsubscribeFromEvents(sub) + + m := NewManager(recorder, statemanager.New(tmpFile)) + m.update = &versionUpdateMock{latestVersion: v.Must(v.NewSemver("1.0.1"))} + m.currentVersion = "1.0.0" + m.autoUpdateSupported = func() bool { return true } + m.Start(context.Background()) + defer m.Stop() + + m.SetVersion("1.0.1", false) + if ver, enforced := waitForUpdateEvent(sub, 500*time.Millisecond); ver != "1.0.1" || !enforced { + t.Fatalf("expected enforced event for 1.0.1, got %q enforced=%v", ver, enforced) + } + + m.SetDownloadOnly() + if ver, enforced := waitForUpdateEvent(sub, 500*time.Millisecond); ver != "1.0.1" || enforced { + t.Fatalf("expected download-only event for 1.0.1, got %q enforced=%v", ver, enforced) + } + + if err := m.Install(context.Background()); err == nil { + t.Fatal("Install in download-only mode must not install the staged managed version") + } +} + +func Test_ResetMode_SilencesStaleForceDirective(t *testing.T) { + tmpFile := path.Join(t.TempDir(), "update-test-force-reset.json") + recorder := peer.NewRecorder("") + sub := recorder.SubscribeToEvents() + defer recorder.UnsubscribeFromEvents(sub) + + mockUpdate := &versionUpdateMock{} + m := NewManager(recorder, statemanager.New(tmpFile)) + m.update = mockUpdate + m.currentVersion = "1.0.0" + m.autoUpdateSupported = func() bool { return true } + m.Start(context.Background()) + defer m.Stop() + + // Management enforces "latest" before the fetcher has reported any version, + // so nothing can be installed while the engine is still up. + m.SetVersion(latestVersion, true) + if event := waitForAnyEvent(sub, 300*time.Millisecond); event != nil { + t.Fatalf("no event expected before the latest version is known, got %v", event) + } + + // The engine stop resets the mode. A release published afterwards must not + // trigger the stale forced install or any notification. + m.ResetMode() + mockUpdate.setLatestVersion(v.Must(v.NewSemver("1.0.1"))) + mockUpdate.onUpdate() + if event := waitForAnyEvent(sub, 300*time.Millisecond); event != nil { + t.Fatalf("stale force directive must stay silent after reset, got %v", event) + } + + m.SetVersion("1.0.1", false) + ver, enforced := waitForUpdateEvent(sub, 500*time.Millisecond) + if ver != "1.0.1" || !enforced { + t.Fatalf("expected enforced event after a fresh directive, got %q enforced=%v", ver, enforced) + } +} + +func Test_SetVersion_MalformedFallsBackToDownloadOnly(t *testing.T) { + tmpFile := path.Join(t.TempDir(), "update-test-malformed.json") + recorder := peer.NewRecorder("") + sub := recorder.SubscribeToEvents() + defer recorder.UnsubscribeFromEvents(sub) + + m := NewManager(recorder, statemanager.New(tmpFile)) + m.update = &versionUpdateMock{latestVersion: v.Must(v.NewSemver("1.0.1"))} + m.currentVersion = "1.0.0" + m.autoUpdateSupported = func() bool { return true } + m.Start(context.Background()) + defer m.Stop() + + m.SetVersion("not-a-version", false) + ver, enforced := waitForUpdateEvent(sub, 500*time.Millisecond) + if ver != "1.0.1" { + t.Fatalf("expected download-only event for 1.0.1 after malformed version, got %q", ver) + } + if enforced { + t.Error("malformed version fallback must not carry enforced metadata") + } +} + +func Test_SetVersion_ForceChangeAppliesWithSameVersion(t *testing.T) { + m := NewManager(peer.NewRecorder(""), statemanager.New(path.Join(t.TempDir(), "update-test-force-change.json"))) + m.update = &versionUpdateMock{} + m.autoUpdateSupported = func() bool { return true } + + m.SetVersion("1.0.1", false) + m.SetVersion("1.0.1", true) + + m.updateMutex.Lock() + defer m.updateMutex.Unlock() + if !m.forceUpdate { + t.Fatal("enabling force update without a version change must take effect") + } + if m.expectedVersion == nil || m.expectedVersion.String() != "1.0.1" { + t.Fatalf("expected version 1.0.1 to stay set, got %v", m.expectedVersion) + } +} + +func Test_SetVersion_RepeatedDirectiveKeepsMode(t *testing.T) { + m := NewManager(peer.NewRecorder(""), statemanager.New(path.Join(t.TempDir(), "update-test-repeat.json"))) + m.update = &versionUpdateMock{} + m.autoUpdateSupported = func() bool { return true } + + for _, expected := range []string{"1.0.1", latestVersion} { + m.SetVersion(expected, false) + m.updateMutex.Lock() + gen := m.modeGen + m.updateMutex.Unlock() + + m.SetVersion(expected, false) + m.updateMutex.Lock() + repeatedGen := m.modeGen + m.updateMutex.Unlock() + + if repeatedGen != gen { + t.Errorf("repeating the %q directive must not reset the mode", expected) + } + } +} diff --git a/client/internal/updater/manager_test.go b/client/internal/updater/manager_test.go index 107dca2b3..939c09814 100644 --- a/client/internal/updater/manager_test.go +++ b/client/internal/updater/manager_test.go @@ -66,7 +66,7 @@ func Test_LatestVersion(t *testing.T) { t.Errorf("%s: Initial update version mismatch, expected %v, got %v", c.name, c.initialLatestVersion.String(), ver) } - mockUpdate.latestVersion = c.latestVersion + mockUpdate.setLatestVersion(c.latestVersion) mockUpdate.onUpdate() ver, _ = waitForUpdateEvent(sub, 500*time.Millisecond) diff --git a/client/internal/updater/manager_test_helpers_test.go b/client/internal/updater/manager_test_helpers_test.go index c7faee1f4..430b31ec1 100644 --- a/client/internal/updater/manager_test_helpers_test.go +++ b/client/internal/updater/manager_test_helpers_test.go @@ -2,21 +2,24 @@ package updater import ( "strconv" + "sync" "time" v "github.com/hashicorp/go-version" "github.com/netbirdio/netbird/client/internal/peer" + cProto "github.com/netbirdio/netbird/client/proto" ) type versionUpdateMock struct { latestVersion *v.Version onUpdate func() + mu sync.Mutex } -func (m versionUpdateMock) StopWatch() {} +func (m *versionUpdateMock) StopWatch() {} -func (m versionUpdateMock) SetDaemonVersion(newVersion string) bool { +func (m *versionUpdateMock) SetDaemonVersion(newVersion string) bool { return false } @@ -24,11 +27,19 @@ func (m *versionUpdateMock) SetOnUpdateListener(updateFn func()) { m.onUpdate = updateFn } -func (m versionUpdateMock) LatestVersion() *v.Version { +func (m *versionUpdateMock) LatestVersion() *v.Version { + m.mu.Lock() + defer m.mu.Unlock() return m.latestVersion } -func (m versionUpdateMock) StartFetcher() {} +func (m *versionUpdateMock) StartFetcher() {} + +func (m *versionUpdateMock) setLatestVersion(version *v.Version) { + m.mu.Lock() + defer m.mu.Unlock() + m.latestVersion = version +} // waitForUpdateEvent waits for a new_version_available event, returns the version string or "" on timeout. func waitForUpdateEvent(sub *peer.EventSubscription, timeout time.Duration) (version string, enforced bool) { @@ -54,3 +65,20 @@ func waitForUpdateEvent(sub *peer.EventSubscription, timeout time.Duration) (ver } } } + +// waitForAnyEvent returns the first published event of any kind, or nil on timeout. +// Unlike waitForUpdateEvent it also catches the install-progress events, so a test +// can assert that a forced install never started. +func waitForAnyEvent(sub *peer.EventSubscription, timeout time.Duration) *cProto.SystemEvent { + timer := time.NewTimer(timeout) + defer timer.Stop() + select { + case event, ok := <-sub.Events(): + if !ok { + return nil + } + return event + case <-timer.C: + return nil + } +}