diff --git a/client/internal/updater/manager.go b/client/internal/updater/manager.go index 48b0cb314..c730713d3 100644 --- a/client/internal/updater/manager.go +++ b/client/internal/updater/manager.go @@ -198,12 +198,13 @@ func (m *Manager) SetVersion(expectedVersion string, forceUpdate bool) { } return } - if m.mode == modeManaged && m.expectedVersion != nil && m.expectedVersion.Equal(parsed) { - return - } expectedSemVer = parsed } + if m.sameDirectiveLocked(expectedSemVer, forceUpdate) { + return + } + m.setModeLocked(modeManaged) m.expectedVersion = expectedSemVer m.updateToLatestVersion = expectedSemVer == nil @@ -444,6 +445,16 @@ func (m *Manager) modeChanged(gen uint64) bool { 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++ diff --git a/client/internal/updater/manager_mode_test.go b/client/internal/updater/manager_mode_test.go index 6a6c52e43..aa238f94c 100644 --- a/client/internal/updater/manager_mode_test.go +++ b/client/internal/updater/manager_mode_test.go @@ -175,3 +175,43 @@ func Test_SetVersion_MalformedFallsBackToDownloadOnly(t *testing.T) { 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) + } + } +}