diff --git a/cmd/cli/commands.go b/cmd/cli/commands.go index ec7b47c..b32f6db 100644 --- a/cmd/cli/commands.go +++ b/cmd/cli/commands.go @@ -1501,28 +1501,31 @@ func initUpgradeCmd() *cobra.Command { if doRestart() { _ = os.Remove(oldBin) _ = os.Chmod(bin, 0755) - ver := "unknown version" - out, err := exec.Command(bin, "--version").CombinedOutput() + ver, err := binaryVersion(bin) if err != nil { mainLog.Load().Warn().Err(err).Msg("Failed to get new binary version") - } - if after, found := strings.CutPrefix(string(out), "ctrld version "); found { - ver = after + ver = "unknown version" } mainLog.Load().Notice().Msgf("Upgrade successful - %s", ver) return } - mainLog.Load().Warn().Msgf("Upgrade failed, restoring previous binary: %s", oldBin) - if err := os.Remove(bin); err != nil { - mainLog.Load().Fatal().Err(err).Msg("failed to remove new binary") + mainLog.Load().Warn().Msg("Upgrade failed: the new binary did not become ready") + stop := func() error { + if !svcInstalled { + return nil + } + if err := stopServiceAndWait(s, upgradeStopTimeout); err != nil { + return err + } + // Mirror the Cleanup task in doRestart: leave DNS settings as the OS + // had them, not as a half-started ctrld left them. + p.router.Cleanup() + p.resetDNS(false, true) + return nil } - if err := os.Rename(oldBin, bin); err != nil { - mainLog.Load().Fatal().Err(err).Msg("failed to restore old binary") - } - if doRestart() { - mainLog.Load().Notice().Msg("Restored previous binary successfully") - return + if err := rollbackToPreviousBinary(bin, oldBin, stop, doRestart); err != nil { + mainLog.Load().Error().Err(err).Msg("Rollback did not complete") } }, } diff --git a/cmd/cli/main_test.go b/cmd/cli/main_test.go index 3e9cd72..6e2257f 100644 --- a/cmd/cli/main_test.go +++ b/cmd/cli/main_test.go @@ -1,6 +1,7 @@ package cli import ( + "fmt" "os" "os/exec" "strings" @@ -11,7 +12,27 @@ import ( var logOutput strings.Builder +// envFakeVersionOutput makes this test binary impersonate a ctrld executable: when +// set, the process writes the value to stdout and exits without running any test, so +// binaryVersion() can be exercised on every platform without building or shipping a +// fixture binary. The value envFakeVersionSilent produces no output at all, which +// reproduces a ctrld.exe_previous that exists but reports no version. +// +// This must be handled before m.Run(), which is what parses the test flags: the child +// is invoked as " --version" and would otherwise die on an unknown flag. +const ( + envFakeVersionOutput = "CTRLD_TEST_FAKE_VERSION_OUTPUT" + envFakeVersionSilent = "" +) + func TestMain(m *testing.M) { + if out := os.Getenv(envFakeVersionOutput); out != "" { + if out != envFakeVersionSilent { + fmt.Println(out) + } + os.Exit(0) + } + l := zerolog.New(&logOutput) mainLog.Store(&l) diff --git a/cmd/cli/upgrade_rollback.go b/cmd/cli/upgrade_rollback.go new file mode 100644 index 0000000..ecda5cf --- /dev/null +++ b/cmd/cli/upgrade_rollback.go @@ -0,0 +1,195 @@ +package cli + +import ( + "context" + "errors" + "fmt" + "os" + "os/exec" + "strings" + "time" + + "github.com/kardianos/service" +) + +const ( + // upgradeStopTimeout bounds how long rollback waits for the replacement process to + // exit, and for Windows to release the lock on its image afterwards. + upgradeStopTimeout = 30 * time.Second + // upgradeStopPollInterval is how often the service status is re-checked while + // waiting for the process to exit. + upgradeStopPollInterval = 500 * time.Millisecond + // binaryVersionTimeout bounds the "--version" probe, so a binary that hangs on + // startup cannot hang the upgrade. + binaryVersionTimeout = 10 * time.Second +) + +// rollbackToPreviousBinary restores oldBin over bin after the replacement failed to +// become ready, and restarts the service on the restored binary. +// +// stop must leave the replacement's process gone, because every step here modifies +// the executable that process is running from. It is called first for that reason: +// readiness failing does not mean the process exited - the service manager can report +// a started service whose process never became operational. Windows holds an +// exclusive lock on a running executable's image, so the previous code's +// os.Remove(bin) failed there with "Access is denied", and because that was fatal the +// restore never ran: the broken binary stayed installed with the previous one +// stranded at its _previous name. +// +// Stopping first also puts the host back in a known state, since a stopped ctrld +// holds no DNS or intercept enforcement. +func rollbackToPreviousBinary(bin, oldBin string, stop func() error, restart func() bool) error { + if err := stop(); err != nil { + mainLog.Load().Error().Err(err).Msg("Could not confirm the service stopped; not modifying its binary") + return err + } + + // Only restore a previous binary that actually runs: a _previous file that exists + // but reports no version would replace a service that starts and hangs with one + // that cannot start at all. + // + // The probe is retried for the same reason removeBinaryWithRetry is: on Windows a + // single exec can fail transiently while antivirus scans the file or the disk is + // busy, and treating that as "no usable previous binary" leaves the host stopped + // with the broken binary installed - an end state worse than restoring a binary + // that turns out to be bad, which the restart check below catches. + // + // Running "--version" proves the file executes. It is not an authenticity check: + // nothing here compares a signature or checksum before a file becomes the installed + // service binary. That is acceptable only because the install directory is writable + // by administrators alone, which is this command's standing assumption. + prevVer, err := binaryVersionWithRetry(oldBin, upgradeStopTimeout) + if err != nil { + mainLog.Load().Error().Err(err).Msgf("Previous binary at %s is not usable, keeping it for inspection", oldBin) + mainLog.Load().Notice().Msgf("Service is stopped and %s is still the installed binary", bin) + return fmt.Errorf("upgrade failed and no usable previous binary to restore: %w", err) + } + + mainLog.Load().Warn().Msgf("Restoring previous binary: %s (%s)", oldBin, prevVer) + if err := removeBinaryWithRetry(bin, upgradeStopTimeout); err != nil { + mainLog.Load().Error().Err(err).Msg("Failed to remove new binary") + mainLog.Load().Notice().Msg("Service is stopped") + return err + } + if err := os.Rename(oldBin, bin); err != nil { + mainLog.Load().Error().Err(err).Msg("Failed to restore old binary") + mainLog.Load().Notice().Msgf("Service is stopped and %s is missing; reinstall ctrld to recover", bin) + return err + } + if restart() { + mainLog.Load().Notice().Msgf("Restored previous binary successfully - %s", prevVer) + return nil + } + + mainLog.Load().Error().Msg("Restored the previous binary but it did not become ready either") + return errors.New("upgrade failed and the restored binary did not become ready") +} + +// stopServiceAndWait stops the service and waits until the service manager reports +// it stopped. Rollback needs the process gone, not merely asked to stop: a stop +// request returns before the process exits, and on Windows the executable stays +// locked until it does. +func stopServiceAndWait(s service.Service, timeout time.Duration) error { + if err := s.Stop(); err != nil { + // Not fatal: the service may already be stopped, or stopping may fail while + // the process is exiting anyway. The status poll below decides. + mainLog.Load().Debug().Err(err).Msg("Stop request failed, waiting for the process to exit anyway") + } + deadline := time.Now().Add(timeout) + statusReadable := false + var lastErr error + for { + status, err := s.Status() + switch { + case errors.Is(err, service.ErrNotInstalled): + return nil + case err == nil: + statusReadable = true + if status == service.StatusStopped { + return nil + } + default: + lastErr = err + } + if !time.Now().Before(deadline) { + if !statusReadable { + // The status was never readable, so "did not stop" was never observed - + // only "could not be observed". Refusing to continue here would leave the + // broken binary installed with the service stopped, which is the outcome + // rollback exists to avoid. Let the caller proceed: the remove is retried + // while the image is locked, and the restart check still has to pass + // before this reports success. + mainLog.Load().Warn().Err(lastErr).Msgf("Could not read service status within %s; continuing with rollback", timeout) + return nil + } + return fmt.Errorf("service did not stop within %s", timeout) + } + time.Sleep(upgradeStopPollInterval) + } +} + +// binaryVersionWithRetry probes a binary's version, retrying transient exec failures +// until timeout. Only the last error is reported: the earlier attempts are noise once a +// retry has been made. +func binaryVersionWithRetry(path string, timeout time.Duration) (string, error) { + deadline := time.Now().Add(timeout) + for { + version, err := binaryVersionFn(path) + if err == nil { + return version, nil + } + if !time.Now().Before(deadline) { + return "", err + } + mainLog.Load().Debug().Err(err).Msgf("Version probe of %s failed, retrying", path) + time.Sleep(upgradeStopPollInterval) + } +} + +// removeBinaryWithRetry removes path, retrying while it is still locked. Windows +// releases the lock on an executable's image asynchronously after its process exits, +// so a remove issued immediately after the service reports stopped can still fail +// with "Access is denied". +func removeBinaryWithRetry(path string, timeout time.Duration) error { + deadline := time.Now().Add(timeout) + for { + err := os.Remove(path) + if err == nil || errors.Is(err, os.ErrNotExist) { + return nil + } + if !time.Now().Before(deadline) { + return fmt.Errorf("could not remove %s within %s: %w", path, timeout, err) + } + time.Sleep(upgradeStopPollInterval) + } +} + +// binaryVersionFn is indirected so rollback can be tested without staging a runnable +// executable per platform. The probe itself is covered directly against the test +// binary; see TestBinaryVersion. +var binaryVersionFn = binaryVersion + +// binaryVersion runs path with "--version" and returns the version it reports. It +// answers "can this binary actually run on this host", which is what rollback needs +// to know before making a file the installed ctrld. +// +// On Windows path is ctrld.exe_previous, whose extension is not in PATHEXT. That +// resolves because os/exec only falls back to appending PATHEXT entries when the path +// has no extension at all (lp_windows.go findExecutable): with one present and the +// file on disk, it is used as-is. A suffix that left no extension - renaming +// oldBinSuffix such that the result is "ctrld_previous" - would break this probe with +// "executable file not found in %PATH%", and rollback would then refuse to restore a +// perfectly good binary. +func binaryVersion(path string) (string, error) { + ctx, cancel := context.WithTimeout(context.Background(), binaryVersionTimeout) + defer cancel() + out, err := exec.CommandContext(ctx, path, "--version").CombinedOutput() + if err != nil { + return "", fmt.Errorf("running %s --version: %w", path, err) + } + ver, found := strings.CutPrefix(strings.TrimSpace(string(out)), "ctrld version ") + if !found { + return "", fmt.Errorf("unexpected --version output from %s: %q", path, strings.TrimSpace(string(out))) + } + return ver, nil +} diff --git a/cmd/cli/upgrade_rollback_test.go b/cmd/cli/upgrade_rollback_test.go new file mode 100644 index 0000000..0c2f169 --- /dev/null +++ b/cmd/cli/upgrade_rollback_test.go @@ -0,0 +1,288 @@ +package cli + +import ( + "errors" + "os" + "path/filepath" + "testing" + "time" + + "github.com/kardianos/service" +) + +// fakeService implements the parts of service.Service that rollback uses. Any other +// method panics, which keeps accidental dependencies visible. +type fakeService struct { + service.Service + + stopErr error + stopCalls int + statuses []service.Status // consumed one per Status() call; the last repeats + statusErr error + onStopCall func() +} + +func (f *fakeService) Stop() error { + f.stopCalls++ + if f.onStopCall != nil { + f.onStopCall() + } + return f.stopErr +} + +func (f *fakeService) Status() (service.Status, error) { + if f.statusErr != nil { + return service.StatusUnknown, f.statusErr + } + if len(f.statuses) == 0 { + return service.StatusStopped, nil + } + st := f.statuses[0] + if len(f.statuses) > 1 { + f.statuses = f.statuses[1:] + } + return st, nil +} + +func TestStopServiceAndWait(t *testing.T) { + tests := []struct { + name string + svc *fakeService + timeout time.Duration + wantErr bool + }{ + { + name: "stops after a few polls", + svc: &fakeService{statuses: []service.Status{service.StatusRunning, service.StatusRunning, service.StatusStopped}}, + timeout: 5 * time.Second, + }, + { + name: "already stopped", + svc: &fakeService{statuses: []service.Status{service.StatusStopped}}, + timeout: 5 * time.Second, + }, + { + // A stop request that errors is not fatal on its own: the process may be + // exiting anyway, so the status poll decides. + name: "stop errors but service is stopped", + svc: &fakeService{stopErr: errors.New("already stopped"), statuses: []service.Status{service.StatusStopped}}, + timeout: 5 * time.Second, + }, + { + name: "not installed", + svc: &fakeService{statusErr: service.ErrNotInstalled}, + timeout: 5 * time.Second, + }, + { + // The process never exits. Rollback must be told so, because modifying a + // running executable is what produced "Access is denied". + name: "never stops", + svc: &fakeService{statuses: []service.Status{service.StatusRunning}}, + timeout: time.Millisecond, + wantErr: true, + }, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + err := stopServiceAndWait(tc.svc, tc.timeout) + if tc.wantErr && err == nil { + t.Fatal("expected an error, got nil") + } + if !tc.wantErr && err != nil { + t.Fatalf("unexpected error: %v", err) + } + if tc.svc.stopCalls != 1 { + t.Errorf("Stop() called %d times, want 1", tc.svc.stopCalls) + } + }) + } +} + +func TestRemoveBinaryWithRetry(t *testing.T) { + t.Run("removes an existing file", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "ctrld") + if err := os.WriteFile(path, []byte("binary"), 0o755); err != nil { + t.Fatal(err) + } + if err := removeBinaryWithRetry(path, time.Second); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if _, err := os.Stat(path); !errors.Is(err, os.ErrNotExist) { + t.Errorf("file still exists after removal: %v", err) + } + }) + + t.Run("missing file is not an error", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "absent") + if err := removeBinaryWithRetry(path, time.Second); err != nil { + t.Fatalf("unexpected error: %v", err) + } + }) + + t.Run("gives up and reports when the path cannot be removed", func(t *testing.T) { + // A non-empty directory stands in for a locked executable: os.Remove keeps + // failing, so the retry loop must surface the error rather than hang. + dir := filepath.Join(t.TempDir(), "locked") + if err := os.Mkdir(dir, 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(dir, "child"), nil, 0o644); err != nil { + t.Fatal(err) + } + if err := removeBinaryWithRetry(dir, time.Millisecond); err == nil { + t.Fatal("expected an error for a path that cannot be removed") + } + }) +} + +func TestBinaryVersion(t *testing.T) { + t.Run("reports the version", func(t *testing.T) { + t.Setenv(envFakeVersionOutput, "ctrld version dev-94fbd3f") + got, err := binaryVersion(os.Args[0]) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != "dev-94fbd3f" { + t.Errorf("binaryVersion() = %q, want %q", got, "dev-94fbd3f") + } + }) + + t.Run("rejects a binary that prints no version", func(t *testing.T) { + // A ctrld.exe_previous that exists and runs, but produces no version output. + // Restoring it would replace a hung service with one that cannot start at all. + t.Setenv(envFakeVersionOutput, envFakeVersionSilent) + if _, err := binaryVersion(os.Args[0]); err == nil { + t.Fatal("expected an error for a binary with no version output") + } + }) + + t.Run("rejects a missing binary", func(t *testing.T) { + if _, err := binaryVersion(filepath.Join(t.TempDir(), "absent")); err == nil { + t.Fatal("expected an error for a missing binary") + } + }) +} + +// stubBinaryVersion makes the version probe report ver for any path, so a rollback +// test does not have to stage a runnable executable. +// +// Staging one is not portable: oldBin is bin+"_previous", so a fixture named "ctrld" +// yields the extension-less "ctrld_previous", which Windows refuses to execute +// ("executable file not found in %PATH%"), and a symlink to the test binary needs a +// privilege Windows does not grant by default. The probe itself is covered against the +// real test binary in TestBinaryVersion; these tests are about rollback's ordering. +func stubBinaryVersion(t *testing.T, ver string, err error) { + t.Helper() + prev := binaryVersionFn + binaryVersionFn = func(string) (string, error) { return ver, err } + t.Cleanup(func() { binaryVersionFn = prev }) +} + +func TestRollbackToPreviousBinaryStopsBeforeTouchingTheBinary(t *testing.T) { + dir := t.TempDir() + bin := filepath.Join(dir, "ctrld") + oldBin := bin + oldBinSuffix + if err := os.WriteFile(bin, []byte("replacement"), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(oldBin, []byte("previous"), 0o755); err != nil { + t.Fatal(err) + } + stubBinaryVersion(t, "dev-a75d669", nil) + + // The invariant: when stop runs, the replacement's executable is still untouched. + // Reversing these two is exactly the "Access is denied" defect. + var stopped bool + var binExistedAtStop bool + stop := func() error { + stopped = true + _, err := os.Stat(bin) + binExistedAtStop = err == nil + return nil + } + restarted := false + restart := func() bool { restarted = true; return true } + + if err := rollbackToPreviousBinary(bin, oldBin, stop, restart); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if !stopped { + t.Error("rollback did not stop the service") + } + if !binExistedAtStop { + t.Error("the binary was modified before the service was stopped") + } + if !restarted { + t.Error("rollback did not restart the service") + } + if _, err := os.Stat(oldBin); !errors.Is(err, os.ErrNotExist) { + t.Errorf("previous binary was not moved into place: %v", err) + } + if _, err := os.Stat(bin); err != nil { + t.Errorf("restored binary is missing: %v", err) + } +} + +func TestRollbackToPreviousBinaryKeepsUnusablePrevious(t *testing.T) { + dir := t.TempDir() + bin := filepath.Join(dir, "ctrld") + oldBin := bin + oldBinSuffix + if err := os.WriteFile(bin, []byte("replacement"), 0o755); err != nil { + t.Fatal(err) + } + // A previous binary that exists but does not report a version. + if err := os.WriteFile(oldBin, []byte("not a working binary"), 0o755); err != nil { + t.Fatal(err) + } + // Stubbed rather than left to the real probe: that would fail here for the right + // reason on unix (not an executable) but the wrong one on Windows (the fixture's + // name has no extension), so the assertion would not be about usability at all. + stubBinaryVersion(t, "", errors.New("unexpected --version output")) + + stopped := false + restarted := false + err := rollbackToPreviousBinary(bin, oldBin, + func() error { stopped = true; return nil }, + func() bool { restarted = true; return true }, + ) + if err == nil { + t.Fatal("expected an error when the previous binary is unusable") + } + if !stopped { + t.Error("the service must still be stopped: a broken replacement holds enforcement") + } + if restarted { + t.Error("must not restart the service with an unusable binary") + } + // Nothing was swapped, and the previous file is kept for inspection. + if _, err := os.Stat(oldBin); err != nil { + t.Errorf("unusable previous binary was not preserved: %v", err) + } + if _, err := os.Stat(bin); err != nil { + t.Errorf("installed binary was removed despite having nothing to restore: %v", err) + } +} + +func TestRollbackToPreviousBinaryAbortsWhenStopFails(t *testing.T) { + dir := t.TempDir() + bin := filepath.Join(dir, "ctrld") + oldBin := bin + oldBinSuffix + for _, p := range []string{bin, oldBin} { + if err := os.WriteFile(p, []byte("binary"), 0o755); err != nil { + t.Fatal(err) + } + } + + stopErr := errors.New("service did not stop within 30s") + err := rollbackToPreviousBinary(bin, oldBin, + func() error { return stopErr }, + func() bool { t.Error("must not restart after a failed stop"); return false }, + ) + if !errors.Is(err, stopErr) { + t.Fatalf("error = %v, want %v", err, stopErr) + } + // The executable of a process that may still be running must be left alone. + if _, err := os.Stat(bin); err != nil { + t.Errorf("binary was modified even though the stop failed: %v", err) + } +}