331 lines
10 KiB
Go
331 lines
10 KiB
Go
package installer
|
|
|
|
import (
|
|
"errors"
|
|
"os"
|
|
"path/filepath"
|
|
"testing"
|
|
)
|
|
|
|
var errSimulatedCrash = errors.New("simulated crash")
|
|
|
|
func TestSwitcherCommitsHealthyUpdate(t *testing.T) {
|
|
root := makeInstallRoot(t, "old", "new")
|
|
switcher := NewSwitcher(func(currentPath string) error {
|
|
assertVersion(t, currentPath, "new")
|
|
return nil
|
|
})
|
|
|
|
if err := switcher.Switch(root); err != nil {
|
|
t.Fatalf("Switch() error = %v", err)
|
|
}
|
|
assertVersion(t, filepath.Join(root, "current"), "new")
|
|
assertMissing(t, filepath.Join(root, "staging"))
|
|
assertMissing(t, filepath.Join(root, "backup"))
|
|
assertMissing(t, filepath.Join(root, transactionFileName))
|
|
}
|
|
|
|
func TestSwitcherRollsBackFailedUpdate(t *testing.T) {
|
|
root := makeInstallRoot(t, "old", "new")
|
|
healthFailure := errors.New("new version did not start")
|
|
switcher := NewSwitcher(func(string) error {
|
|
return healthFailure
|
|
})
|
|
|
|
err := switcher.Switch(root)
|
|
if !errors.Is(err, ErrHealthCheckFailed) {
|
|
t.Fatalf("Switch() error = %v, want %v", err, ErrHealthCheckFailed)
|
|
}
|
|
assertVersion(t, filepath.Join(root, "current"), "old")
|
|
assertMissing(t, filepath.Join(root, "staging"))
|
|
assertMissing(t, filepath.Join(root, "backup"))
|
|
assertMissing(t, filepath.Join(root, transactionFileName))
|
|
}
|
|
|
|
func TestSwitcherRemovesFailedInitialInstall(t *testing.T) {
|
|
root := makeInstallRoot(t, "", "new")
|
|
switcher := NewSwitcher(func(string) error {
|
|
return errors.New("health failed")
|
|
})
|
|
|
|
err := switcher.Switch(root)
|
|
if !errors.Is(err, ErrHealthCheckFailed) {
|
|
t.Fatalf("Switch() error = %v, want %v", err, ErrHealthCheckFailed)
|
|
}
|
|
assertMissing(t, filepath.Join(root, "current"))
|
|
assertMissing(t, filepath.Join(root, "staging"))
|
|
assertMissing(t, filepath.Join(root, transactionFileName))
|
|
}
|
|
|
|
func TestSwitcherPreSwitchCheckPreventsUpdateAndCleansStaging(t *testing.T) {
|
|
root := makeInstallRoot(t, "old", "new")
|
|
checkErr := errors.New("target is running")
|
|
switcher := NewSwitcherWithPreSwitchCheck(
|
|
func(string) error { return nil },
|
|
func() error {
|
|
assertVersion(t, filepath.Join(root, "current"), "old")
|
|
assertVersion(t, filepath.Join(root, "staging"), "new")
|
|
return checkErr
|
|
},
|
|
)
|
|
|
|
err := switcher.Switch(root)
|
|
if !errors.Is(err, checkErr) {
|
|
t.Fatalf("Switch() error = %v, want %v", err, checkErr)
|
|
}
|
|
assertVersion(t, filepath.Join(root, "current"), "old")
|
|
assertMissing(t, filepath.Join(root, "staging"))
|
|
assertMissing(t, filepath.Join(root, "backup"))
|
|
assertMissing(t, filepath.Join(root, transactionFileName))
|
|
}
|
|
|
|
func TestSwitcherSkipsPreSwitchCheckForInitialInstall(t *testing.T) {
|
|
root := makeInstallRoot(t, "", "new")
|
|
switcher := NewSwitcherWithPreSwitchCheck(
|
|
func(string) error { return nil },
|
|
func() error { return errors.New("must not run") },
|
|
)
|
|
|
|
if err := switcher.Switch(root); err != nil {
|
|
t.Fatalf("Switch() error = %v", err)
|
|
}
|
|
assertVersion(t, filepath.Join(root, "current"), "new")
|
|
}
|
|
|
|
func TestRecoverInterruptedSwitch(t *testing.T) {
|
|
tests := []struct {
|
|
name string
|
|
crashStep switchStep
|
|
wantVersion string
|
|
wantAction RecoveryAction
|
|
}{
|
|
{
|
|
name: "after current rename before phase update",
|
|
crashStep: stepCurrentRenamed,
|
|
wantVersion: "old",
|
|
wantAction: RecoveryRolledBack,
|
|
},
|
|
{
|
|
name: "after staging rename before phase update",
|
|
crashStep: stepStagingRenamed,
|
|
wantVersion: "old",
|
|
wantAction: RecoveryRolledBack,
|
|
},
|
|
{
|
|
name: "after committed journal before cleanup",
|
|
crashStep: stepCommitted,
|
|
wantVersion: "new",
|
|
wantAction: RecoveryCommitted,
|
|
},
|
|
}
|
|
|
|
for _, test := range tests {
|
|
t.Run(test.name, func(t *testing.T) {
|
|
root := makeInstallRoot(t, "old", "new")
|
|
switcher := NewSwitcher(func(string) error { return nil })
|
|
switcher.afterStep = func(step switchStep) error {
|
|
if step == test.crashStep {
|
|
return errSimulatedCrash
|
|
}
|
|
return nil
|
|
}
|
|
|
|
err := switcher.Switch(root)
|
|
if !errors.Is(err, errSimulatedCrash) {
|
|
t.Fatalf("Switch() error = %v, want %v", err, errSimulatedCrash)
|
|
}
|
|
|
|
result, err := Recover(root)
|
|
if err != nil {
|
|
t.Fatalf("Recover() error = %v", err)
|
|
}
|
|
if result.Action != test.wantAction {
|
|
t.Fatalf("Action = %q, want %q", result.Action, test.wantAction)
|
|
}
|
|
assertVersion(t, filepath.Join(root, "current"), test.wantVersion)
|
|
assertMissing(t, filepath.Join(root, "staging"))
|
|
assertMissing(t, filepath.Join(root, "backup"))
|
|
assertMissing(t, filepath.Join(root, transactionFileName))
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestRecoverReadsTransactionBackup(t *testing.T) {
|
|
root := makeInstallRoot(t, "old", "new")
|
|
switcher := NewSwitcher(func(string) error { return nil })
|
|
switcher.afterStep = func(step switchStep) error {
|
|
if step == stepStagingRenamed {
|
|
return errSimulatedCrash
|
|
}
|
|
return nil
|
|
}
|
|
if err := switcher.Switch(root); !errors.Is(err, errSimulatedCrash) {
|
|
t.Fatalf("Switch() error = %v, want %v", err, errSimulatedCrash)
|
|
}
|
|
if err := os.Rename(
|
|
filepath.Join(root, transactionFileName),
|
|
filepath.Join(root, transactionBackupFileName),
|
|
); err != nil {
|
|
t.Fatalf("move transaction to backup: %v", err)
|
|
}
|
|
|
|
result, err := Recover(root)
|
|
if err != nil {
|
|
t.Fatalf("Recover() error = %v", err)
|
|
}
|
|
if result.Action != RecoveryRolledBack {
|
|
t.Fatalf("Action = %q, want %q", result.Action, RecoveryRolledBack)
|
|
}
|
|
assertVersion(t, filepath.Join(root, "current"), "old")
|
|
assertMissing(t, filepath.Join(root, transactionBackupFileName))
|
|
}
|
|
|
|
func TestRecoverAbortsInterruptedInitialInstall(t *testing.T) {
|
|
root := makeInstallRoot(t, "", "new")
|
|
switcher := NewSwitcher(func(string) error { return nil })
|
|
switcher.afterStep = func(step switchStep) error {
|
|
if step == stepStagingRenamed {
|
|
return errSimulatedCrash
|
|
}
|
|
return nil
|
|
}
|
|
if err := switcher.Switch(root); !errors.Is(err, errSimulatedCrash) {
|
|
t.Fatalf("Switch() error = %v, want %v", err, errSimulatedCrash)
|
|
}
|
|
|
|
result, err := Recover(root)
|
|
if err != nil {
|
|
t.Fatalf("Recover() error = %v", err)
|
|
}
|
|
if result.Action != RecoveryAborted {
|
|
t.Fatalf("Action = %q, want %q", result.Action, RecoveryAborted)
|
|
}
|
|
assertMissing(t, filepath.Join(root, "current"))
|
|
assertMissing(t, filepath.Join(root, "staging"))
|
|
assertMissing(t, filepath.Join(root, transactionFileName))
|
|
}
|
|
|
|
func TestRecoverCompletesInterruptedRollback(t *testing.T) {
|
|
root := makeInstallRoot(t, "old", "new")
|
|
switcher := NewSwitcher(func(string) error {
|
|
return errors.New("health failed")
|
|
})
|
|
switcher.afterStep = func(step switchStep) error {
|
|
if step == stepRollbackRequired {
|
|
return errSimulatedCrash
|
|
}
|
|
return nil
|
|
}
|
|
if err := switcher.Switch(root); !errors.Is(err, errSimulatedCrash) {
|
|
t.Fatalf("Switch() error = %v, want %v", err, errSimulatedCrash)
|
|
}
|
|
|
|
result, err := Recover(root)
|
|
if err != nil {
|
|
t.Fatalf("Recover() error = %v", err)
|
|
}
|
|
if result.Action != RecoveryRolledBack {
|
|
t.Fatalf("Action = %q, want %q", result.Action, RecoveryRolledBack)
|
|
}
|
|
assertVersion(t, filepath.Join(root, "current"), "old")
|
|
assertMissing(t, filepath.Join(root, "staging"))
|
|
assertMissing(t, filepath.Join(root, "backup"))
|
|
}
|
|
|
|
func TestSwitcherRejectsUnsafeStartingState(t *testing.T) {
|
|
t.Run("missing staging", func(t *testing.T) {
|
|
root := makeInstallRoot(t, "old", "")
|
|
err := NewSwitcher(func(string) error { return nil }).Switch(root)
|
|
if !errors.Is(err, ErrStagingMissing) {
|
|
t.Fatalf("Switch() error = %v, want %v", err, ErrStagingMissing)
|
|
}
|
|
})
|
|
|
|
t.Run("existing backup", func(t *testing.T) {
|
|
root := makeInstallRoot(t, "old", "new")
|
|
writeVersion(t, filepath.Join(root, "backup"), "stale")
|
|
err := NewSwitcher(func(string) error { return nil }).Switch(root)
|
|
if !errors.Is(err, ErrBackupExists) {
|
|
t.Fatalf("Switch() error = %v, want %v", err, ErrBackupExists)
|
|
}
|
|
})
|
|
|
|
t.Run("staging symlink", func(t *testing.T) {
|
|
root := t.TempDir()
|
|
target := filepath.Join(t.TempDir(), "target")
|
|
writeVersion(t, target, "new")
|
|
if err := os.Symlink(target, filepath.Join(root, "staging")); err != nil {
|
|
t.Skipf("symlink unavailable: %v", err)
|
|
}
|
|
err := NewSwitcher(func(string) error { return nil }).Switch(root)
|
|
if !errors.Is(err, ErrUnsafeInstallLayout) {
|
|
t.Fatalf("Switch() error = %v, want %v", err, ErrUnsafeInstallLayout)
|
|
}
|
|
})
|
|
|
|
t.Run("transaction symlink", func(t *testing.T) {
|
|
root := makeInstallRoot(t, "old", "new")
|
|
target := filepath.Join(t.TempDir(), "transaction.json")
|
|
if err := os.WriteFile(target, []byte(`{}`), 0o600); err != nil {
|
|
t.Fatalf("write symlink target: %v", err)
|
|
}
|
|
if err := os.Symlink(target, filepath.Join(root, transactionFileName)); err != nil {
|
|
t.Skipf("symlink unavailable: %v", err)
|
|
}
|
|
err := NewSwitcher(func(string) error { return nil }).Switch(root)
|
|
if !errors.Is(err, ErrUnsafeInstallLayout) {
|
|
t.Fatalf("Switch() error = %v, want %v", err, ErrUnsafeInstallLayout)
|
|
}
|
|
})
|
|
}
|
|
|
|
func TestRecoverRejectsOrphanBackup(t *testing.T) {
|
|
root := makeInstallRoot(t, "old", "")
|
|
writeVersion(t, filepath.Join(root, "backup"), "orphan")
|
|
|
|
_, err := Recover(root)
|
|
if !errors.Is(err, ErrRecoveryInconsistent) {
|
|
t.Fatalf("Recover() error = %v, want %v", err, ErrRecoveryInconsistent)
|
|
}
|
|
}
|
|
|
|
func makeInstallRoot(t *testing.T, currentVersion, stagingVersion string) string {
|
|
t.Helper()
|
|
root := t.TempDir()
|
|
if currentVersion != "" {
|
|
writeVersion(t, filepath.Join(root, "current"), currentVersion)
|
|
}
|
|
if stagingVersion != "" {
|
|
writeVersion(t, filepath.Join(root, "staging"), stagingVersion)
|
|
}
|
|
return root
|
|
}
|
|
|
|
func writeVersion(t *testing.T, directory, version string) {
|
|
t.Helper()
|
|
if err := os.MkdirAll(directory, 0o700); err != nil {
|
|
t.Fatalf("create version directory: %v", err)
|
|
}
|
|
if err := os.WriteFile(filepath.Join(directory, "version.txt"), []byte(version), 0o600); err != nil {
|
|
t.Fatalf("write version: %v", err)
|
|
}
|
|
}
|
|
|
|
func assertVersion(t *testing.T, directory, want string) {
|
|
t.Helper()
|
|
data, err := os.ReadFile(filepath.Join(directory, "version.txt"))
|
|
if err != nil {
|
|
t.Fatalf("read version from %s: %v", directory, err)
|
|
}
|
|
if string(data) != want {
|
|
t.Fatalf("version in %s = %q, want %q", directory, data, want)
|
|
}
|
|
}
|
|
|
|
func assertMissing(t *testing.T, path string) {
|
|
t.Helper()
|
|
if _, err := os.Stat(path); !os.IsNotExist(err) {
|
|
t.Fatalf("%s should be missing, stat error = %v", path, err)
|
|
}
|
|
}
|