Prototype atomic install recovery (T-103)
This commit is contained in:
@@ -0,0 +1,110 @@
|
||||
package installer
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrUnsafeInstallLayout = errors.New("unsafe install directory layout")
|
||||
ErrStagingMissing = errors.New("staging directory is missing")
|
||||
ErrBackupExists = errors.New("backup directory already exists")
|
||||
)
|
||||
|
||||
type appLayout struct {
|
||||
root string
|
||||
current string
|
||||
staging string
|
||||
backup string
|
||||
transaction string
|
||||
transactionBackup string
|
||||
}
|
||||
|
||||
type directoryState struct {
|
||||
current bool
|
||||
staging bool
|
||||
backup bool
|
||||
}
|
||||
|
||||
func inspectAppLayout(root string) (appLayout, error) {
|
||||
if root == "" {
|
||||
return appLayout{}, fmt.Errorf("%w: empty root", ErrUnsafeInstallLayout)
|
||||
}
|
||||
absolute, err := filepath.Abs(root)
|
||||
if err != nil {
|
||||
return appLayout{}, fmt.Errorf("%w: %v", ErrUnsafeInstallLayout, err)
|
||||
}
|
||||
info, err := os.Lstat(absolute)
|
||||
if err != nil {
|
||||
return appLayout{}, fmt.Errorf("%w: root: %v", ErrUnsafeInstallLayout, err)
|
||||
}
|
||||
if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() {
|
||||
return appLayout{}, fmt.Errorf("%w: root is not a real directory", ErrUnsafeInstallLayout)
|
||||
}
|
||||
transaction, transactionBackup := transactionPaths(absolute)
|
||||
return appLayout{
|
||||
root: absolute,
|
||||
current: filepath.Join(absolute, "current"),
|
||||
staging: filepath.Join(absolute, "staging"),
|
||||
backup: filepath.Join(absolute, "backup"),
|
||||
transaction: transaction,
|
||||
transactionBackup: transactionBackup,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func inspectDirectories(layout appLayout) (directoryState, error) {
|
||||
current, err := inspectManagedDirectory(layout.current)
|
||||
if err != nil {
|
||||
return directoryState{}, err
|
||||
}
|
||||
staging, err := inspectManagedDirectory(layout.staging)
|
||||
if err != nil {
|
||||
return directoryState{}, err
|
||||
}
|
||||
backup, err := inspectManagedDirectory(layout.backup)
|
||||
if err != nil {
|
||||
return directoryState{}, err
|
||||
}
|
||||
return directoryState{
|
||||
current: current,
|
||||
staging: staging,
|
||||
backup: backup,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func inspectManagedDirectory(path string) (bool, error) {
|
||||
info, err := os.Lstat(path)
|
||||
if os.IsNotExist(err) {
|
||||
return false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("%w: inspect %s: %v", ErrUnsafeInstallLayout, path, err)
|
||||
}
|
||||
if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() {
|
||||
return false, fmt.Errorf("%w: %s is not a real directory", ErrUnsafeInstallLayout, path)
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func removeManagedDirectory(layout appLayout, target string) error {
|
||||
if filepath.Dir(target) != layout.root {
|
||||
return fmt.Errorf("%w: refuse removal outside app root", ErrUnsafeInstallLayout)
|
||||
}
|
||||
base := filepath.Base(target)
|
||||
if base != "current" && base != "staging" && base != "backup" {
|
||||
return fmt.Errorf("%w: refuse removal of %s", ErrUnsafeInstallLayout, base)
|
||||
}
|
||||
exists, err := inspectManagedDirectory(target)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !exists {
|
||||
return nil
|
||||
}
|
||||
if err := os.RemoveAll(target); err != nil {
|
||||
return fmt.Errorf("remove managed directory %s: %w", base, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,173 @@
|
||||
package installer
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
)
|
||||
|
||||
var ErrRecoveryInconsistent = errors.New("install recovery state is inconsistent")
|
||||
|
||||
type RecoveryAction string
|
||||
|
||||
const (
|
||||
RecoveryNone RecoveryAction = "none"
|
||||
RecoveryAborted RecoveryAction = "aborted"
|
||||
RecoveryRolledBack RecoveryAction = "rolled_back"
|
||||
RecoveryCommitted RecoveryAction = "committed"
|
||||
)
|
||||
|
||||
type RecoveryResult struct {
|
||||
Action RecoveryAction
|
||||
Phase string
|
||||
}
|
||||
|
||||
// Recover resolves an interrupted transaction from journal and directory state.
|
||||
func Recover(root string) (RecoveryResult, error) {
|
||||
layout, err := inspectAppLayout(root)
|
||||
if err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
record, exists, err := loadTransaction(layout)
|
||||
if err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
state, err := inspectDirectories(layout)
|
||||
if err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
if !exists {
|
||||
if state.backup {
|
||||
return RecoveryResult{}, fmt.Errorf(
|
||||
"%w: backup exists without transaction",
|
||||
ErrRecoveryInconsistent,
|
||||
)
|
||||
}
|
||||
return RecoveryResult{Action: RecoveryNone}, nil
|
||||
}
|
||||
|
||||
result := RecoveryResult{Phase: string(record.Phase)}
|
||||
switch record.Phase {
|
||||
case phaseCommitted:
|
||||
if !state.current || state.staging {
|
||||
return RecoveryResult{}, fmt.Errorf(
|
||||
"%w: committed current=%t staging=%t",
|
||||
ErrRecoveryInconsistent,
|
||||
state.current,
|
||||
state.staging,
|
||||
)
|
||||
}
|
||||
if err := removeManagedDirectory(layout, layout.backup); err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
if err := removeTransaction(layout); err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
result.Action = RecoveryCommitted
|
||||
return result, nil
|
||||
case phaseRollbackRequired:
|
||||
return recoverRollbackRequired(layout, record, state, result)
|
||||
case phasePrepared, phaseCurrentBackedUp, phaseStagingActivated:
|
||||
return recoverUncommitted(layout, record, state, result)
|
||||
default:
|
||||
return RecoveryResult{}, fmt.Errorf("%w: phase=%q", ErrTransactionCorrupt, record.Phase)
|
||||
}
|
||||
}
|
||||
|
||||
func recoverUncommitted(
|
||||
layout appLayout,
|
||||
record transactionRecord,
|
||||
state directoryState,
|
||||
result RecoveryResult,
|
||||
) (RecoveryResult, error) {
|
||||
if record.HadCurrent {
|
||||
if state.backup {
|
||||
if state.current && state.staging {
|
||||
return RecoveryResult{}, fmt.Errorf(
|
||||
"%w: current, staging and backup all exist",
|
||||
ErrRecoveryInconsistent,
|
||||
)
|
||||
}
|
||||
if state.current {
|
||||
if err := os.Rename(layout.current, layout.staging); err != nil {
|
||||
return RecoveryResult{}, fmt.Errorf("move unverified current aside: %w", err)
|
||||
}
|
||||
}
|
||||
if err := os.Rename(layout.backup, layout.current); err != nil {
|
||||
return RecoveryResult{}, fmt.Errorf("restore backup during recovery: %w", err)
|
||||
}
|
||||
if err := removeManagedDirectory(layout, layout.staging); err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
if err := removeTransaction(layout); err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
result.Action = RecoveryRolledBack
|
||||
return result, nil
|
||||
}
|
||||
if record.Phase == phasePrepared && state.current && state.staging {
|
||||
if err := removeManagedDirectory(layout, layout.staging); err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
if err := removeTransaction(layout); err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
result.Action = RecoveryAborted
|
||||
return result, nil
|
||||
}
|
||||
return RecoveryResult{}, fmt.Errorf(
|
||||
"%w: old current has no recoverable backup",
|
||||
ErrRecoveryInconsistent,
|
||||
)
|
||||
}
|
||||
|
||||
if state.backup || (state.current && state.staging) {
|
||||
return RecoveryResult{}, fmt.Errorf(
|
||||
"%w: initial install has conflicting directories",
|
||||
ErrRecoveryInconsistent,
|
||||
)
|
||||
}
|
||||
if err := removeManagedDirectory(layout, layout.current); err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
if err := removeManagedDirectory(layout, layout.staging); err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
if err := removeTransaction(layout); err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
result.Action = RecoveryAborted
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func recoverRollbackRequired(
|
||||
layout appLayout,
|
||||
record transactionRecord,
|
||||
state directoryState,
|
||||
result RecoveryResult,
|
||||
) (RecoveryResult, error) {
|
||||
if record.HadCurrent && !state.backup {
|
||||
if !state.current {
|
||||
return RecoveryResult{}, fmt.Errorf(
|
||||
"%w: rollback lost current and backup",
|
||||
ErrRecoveryInconsistent,
|
||||
)
|
||||
}
|
||||
if err := removeManagedDirectory(layout, layout.staging); err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
if err := removeTransaction(layout); err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
result.Action = RecoveryRolledBack
|
||||
return result, nil
|
||||
}
|
||||
if err := rollbackActivated(layout, record.HadCurrent); err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
if err := removeTransaction(layout); err != nil {
|
||||
return RecoveryResult{}, err
|
||||
}
|
||||
result.Action = RecoveryRolledBack
|
||||
return result, nil
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
package installer
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrHealthCheckRequired = errors.New("health check is required")
|
||||
ErrHealthCheckFailed = errors.New("installed version failed health check")
|
||||
ErrRollbackFailed = errors.New("install rollback failed")
|
||||
)
|
||||
|
||||
type HealthCheck func(currentPath string) error
|
||||
|
||||
type switchStep string
|
||||
|
||||
const (
|
||||
stepPrepared switchStep = "prepared"
|
||||
stepCurrentRenamed switchStep = "current_renamed"
|
||||
stepCurrentBackedUp switchStep = "current_backed_up"
|
||||
stepStagingRenamed switchStep = "staging_renamed"
|
||||
stepStagingActivated switchStep = "staging_activated"
|
||||
stepRollbackRequired switchStep = "rollback_required"
|
||||
stepCommitted switchStep = "committed"
|
||||
)
|
||||
|
||||
// RollbackError reports both the health failure and rollback failure.
|
||||
type RollbackError struct {
|
||||
Health error
|
||||
Rollback error
|
||||
}
|
||||
|
||||
func (err *RollbackError) Error() string {
|
||||
return fmt.Sprintf("%s: health=%v; rollback=%v", ErrRollbackFailed, err.Health, err.Rollback)
|
||||
}
|
||||
|
||||
func (err *RollbackError) Unwrap() error {
|
||||
return ErrRollbackFailed
|
||||
}
|
||||
|
||||
// Switcher activates a verified staging directory and runs an injected check.
|
||||
type Switcher struct {
|
||||
health HealthCheck
|
||||
afterStep func(switchStep) error
|
||||
}
|
||||
|
||||
func NewSwitcher(health HealthCheck) *Switcher {
|
||||
return &Switcher{health: health}
|
||||
}
|
||||
|
||||
func (switcher *Switcher) Switch(root string) error {
|
||||
if switcher.health == nil {
|
||||
return ErrHealthCheckRequired
|
||||
}
|
||||
layout, err := inspectAppLayout(root)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, exists, err := loadTransaction(layout); err != nil {
|
||||
return err
|
||||
} else if exists {
|
||||
return ErrRecoveryRequired
|
||||
}
|
||||
state, err := inspectDirectories(layout)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !state.staging {
|
||||
return ErrStagingMissing
|
||||
}
|
||||
if state.backup {
|
||||
return ErrBackupExists
|
||||
}
|
||||
|
||||
record := newTransaction(phasePrepared, state.current)
|
||||
if err := writeTransaction(layout, record); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := switcher.runStep(stepPrepared); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if state.current {
|
||||
if err := os.Rename(layout.current, layout.backup); err != nil {
|
||||
return fmt.Errorf("backup current directory: %w", err)
|
||||
}
|
||||
if err := switcher.runStep(stepCurrentRenamed); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
record.Phase = phaseCurrentBackedUp
|
||||
if err := writeTransaction(layout, record); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := switcher.runStep(stepCurrentBackedUp); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := os.Rename(layout.staging, layout.current); err != nil {
|
||||
return fmt.Errorf("activate staging directory: %w", err)
|
||||
}
|
||||
if err := switcher.runStep(stepStagingRenamed); err != nil {
|
||||
return err
|
||||
}
|
||||
record.Phase = phaseStagingActivated
|
||||
if err := writeTransaction(layout, record); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := switcher.runStep(stepStagingActivated); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
healthErr := switcher.health(layout.current)
|
||||
if healthErr != nil {
|
||||
record.Phase = phaseRollbackRequired
|
||||
if err := writeTransaction(layout, record); err != nil {
|
||||
return &RollbackError{Health: healthErr, Rollback: err}
|
||||
}
|
||||
if err := switcher.runStep(stepRollbackRequired); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := rollbackActivated(layout, record.HadCurrent); err != nil {
|
||||
return &RollbackError{Health: healthErr, Rollback: err}
|
||||
}
|
||||
if err := removeTransaction(layout); err != nil {
|
||||
return &RollbackError{Health: healthErr, Rollback: err}
|
||||
}
|
||||
return fmt.Errorf("%w: %v", ErrHealthCheckFailed, healthErr)
|
||||
}
|
||||
|
||||
record.Phase = phaseCommitted
|
||||
if err := writeTransaction(layout, record); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := switcher.runStep(stepCommitted); err != nil {
|
||||
return err
|
||||
}
|
||||
if record.HadCurrent {
|
||||
if err := removeManagedDirectory(layout, layout.backup); err != nil {
|
||||
return fmt.Errorf("%w: cleanup committed backup: %v", ErrRecoveryRequired, err)
|
||||
}
|
||||
}
|
||||
if err := removeTransaction(layout); err != nil {
|
||||
return fmt.Errorf("%w: %v", ErrRecoveryRequired, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (switcher *Switcher) runStep(step switchStep) error {
|
||||
if switcher.afterStep == nil {
|
||||
return nil
|
||||
}
|
||||
return switcher.afterStep(step)
|
||||
}
|
||||
|
||||
func rollbackActivated(layout appLayout, hadCurrent bool) error {
|
||||
state, err := inspectDirectories(layout)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if hadCurrent {
|
||||
if !state.backup {
|
||||
return fmt.Errorf("%w: previous current backup is missing", ErrRollbackFailed)
|
||||
}
|
||||
if state.current {
|
||||
if state.staging {
|
||||
return fmt.Errorf("%w: current and staging both exist", ErrRollbackFailed)
|
||||
}
|
||||
if err := os.Rename(layout.current, layout.staging); err != nil {
|
||||
return fmt.Errorf("move failed current aside: %w", err)
|
||||
}
|
||||
}
|
||||
if err := os.Rename(layout.backup, layout.current); err != nil {
|
||||
if _, statErr := os.Stat(layout.staging); statErr == nil {
|
||||
_ = os.Rename(layout.staging, layout.current)
|
||||
}
|
||||
return fmt.Errorf("restore previous current: %w", err)
|
||||
}
|
||||
return removeManagedDirectory(layout, layout.staging)
|
||||
}
|
||||
|
||||
if state.backup {
|
||||
return fmt.Errorf("%w: unexpected backup without previous current", ErrRollbackFailed)
|
||||
}
|
||||
if state.current {
|
||||
if state.staging {
|
||||
return fmt.Errorf("%w: current and staging both exist", ErrRollbackFailed)
|
||||
}
|
||||
if err := os.Rename(layout.current, layout.staging); err != nil {
|
||||
return fmt.Errorf("move failed initial install aside: %w", err)
|
||||
}
|
||||
}
|
||||
return removeManagedDirectory(layout, layout.staging)
|
||||
}
|
||||
@@ -0,0 +1,295 @@
|
||||
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 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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,209 @@
|
||||
package installer
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
)
|
||||
|
||||
const (
|
||||
transactionFileName = "install-transaction.json"
|
||||
transactionBackupFileName = "install-transaction.json.backup"
|
||||
transactionSchemaVersion = 1
|
||||
)
|
||||
|
||||
type transactionPhase string
|
||||
|
||||
const (
|
||||
phasePrepared transactionPhase = "prepared"
|
||||
phaseCurrentBackedUp transactionPhase = "current_backed_up"
|
||||
phaseStagingActivated transactionPhase = "staging_activated"
|
||||
phaseRollbackRequired transactionPhase = "rollback_required"
|
||||
phaseCommitted transactionPhase = "committed"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrTransactionCorrupt = errors.New("install transaction is corrupt")
|
||||
ErrRecoveryRequired = errors.New("install recovery is required")
|
||||
)
|
||||
|
||||
type transactionRecord struct {
|
||||
SchemaVersion int `json:"schema_version"`
|
||||
Phase transactionPhase `json:"phase"`
|
||||
HadCurrent bool `json:"had_current"`
|
||||
}
|
||||
|
||||
func newTransaction(phase transactionPhase, hadCurrent bool) transactionRecord {
|
||||
return transactionRecord{
|
||||
SchemaVersion: transactionSchemaVersion,
|
||||
Phase: phase,
|
||||
HadCurrent: hadCurrent,
|
||||
}
|
||||
}
|
||||
|
||||
func (record transactionRecord) validate() error {
|
||||
if record.SchemaVersion != transactionSchemaVersion {
|
||||
return fmt.Errorf(
|
||||
"%w: schema_version=%d",
|
||||
ErrTransactionCorrupt,
|
||||
record.SchemaVersion,
|
||||
)
|
||||
}
|
||||
switch record.Phase {
|
||||
case phasePrepared,
|
||||
phaseCurrentBackedUp,
|
||||
phaseStagingActivated,
|
||||
phaseRollbackRequired,
|
||||
phaseCommitted:
|
||||
return nil
|
||||
default:
|
||||
return fmt.Errorf("%w: phase=%q", ErrTransactionCorrupt, record.Phase)
|
||||
}
|
||||
}
|
||||
|
||||
func writeTransaction(layout appLayout, record transactionRecord) error {
|
||||
if err := record.validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := json.Marshal(record)
|
||||
if err != nil {
|
||||
return fmt.Errorf("encode install transaction: %w", err)
|
||||
}
|
||||
data = append(data, '\n')
|
||||
if err := replaceFileWithBackup(
|
||||
layout.root,
|
||||
layout.transaction,
|
||||
layout.transactionBackup,
|
||||
data,
|
||||
); err != nil {
|
||||
return fmt.Errorf("write install transaction: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func loadTransaction(layout appLayout) (transactionRecord, bool, error) {
|
||||
data, err := readTransactionFile(layout.transaction)
|
||||
if os.IsNotExist(err) {
|
||||
data, err = readTransactionFile(layout.transactionBackup)
|
||||
}
|
||||
if os.IsNotExist(err) {
|
||||
return transactionRecord{}, false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return transactionRecord{}, false, fmt.Errorf("read install transaction: %w", err)
|
||||
}
|
||||
|
||||
decoder := json.NewDecoder(bytes.NewReader(data))
|
||||
decoder.DisallowUnknownFields()
|
||||
var record transactionRecord
|
||||
if err := decoder.Decode(&record); err != nil {
|
||||
return transactionRecord{}, false, fmt.Errorf("%w: %v", ErrTransactionCorrupt, err)
|
||||
}
|
||||
if err := ensureJSONEOF(decoder); err != nil {
|
||||
return transactionRecord{}, false, err
|
||||
}
|
||||
if err := record.validate(); err != nil {
|
||||
return transactionRecord{}, false, err
|
||||
}
|
||||
return record, true, nil
|
||||
}
|
||||
|
||||
func readTransactionFile(path string) ([]byte, error) {
|
||||
info, err := os.Lstat(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if info.Mode()&os.ModeSymlink != 0 || !info.Mode().IsRegular() {
|
||||
return nil, fmt.Errorf("%w: transaction is not a regular file", ErrUnsafeInstallLayout)
|
||||
}
|
||||
return os.ReadFile(path)
|
||||
}
|
||||
|
||||
func ensureJSONEOF(decoder *json.Decoder) error {
|
||||
var extra any
|
||||
if err := decoder.Decode(&extra); err != io.EOF {
|
||||
if err == nil {
|
||||
return fmt.Errorf("%w: trailing JSON value", ErrTransactionCorrupt)
|
||||
}
|
||||
return fmt.Errorf("%w: trailing data: %v", ErrTransactionCorrupt, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func removeTransaction(layout appLayout) error {
|
||||
if err := os.Remove(layout.transaction); err != nil && !os.IsNotExist(err) {
|
||||
return fmt.Errorf("remove install transaction: %w", err)
|
||||
}
|
||||
if err := os.Remove(layout.transactionBackup); err != nil && !os.IsNotExist(err) {
|
||||
return fmt.Errorf("remove install transaction backup: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func replaceFileWithBackup(
|
||||
directory string,
|
||||
target string,
|
||||
backup string,
|
||||
data []byte,
|
||||
) error {
|
||||
temporary, err := os.CreateTemp(directory, ".install-transaction-*.tmp")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
temporaryPath := temporary.Name()
|
||||
defer os.Remove(temporaryPath)
|
||||
|
||||
if err := temporary.Chmod(0o600); err != nil {
|
||||
temporary.Close()
|
||||
return err
|
||||
}
|
||||
if _, err := temporary.Write(data); err != nil {
|
||||
temporary.Close()
|
||||
return err
|
||||
}
|
||||
if err := temporary.Sync(); err != nil {
|
||||
temporary.Close()
|
||||
return err
|
||||
}
|
||||
if err := temporary.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
movedTarget := false
|
||||
if info, err := os.Lstat(target); err == nil {
|
||||
if info.Mode()&os.ModeSymlink != 0 || !info.Mode().IsRegular() {
|
||||
return fmt.Errorf("%w: transaction target is not a regular file", ErrUnsafeInstallLayout)
|
||||
}
|
||||
if err := os.Remove(backup); err != nil && !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
if err := os.Rename(target, backup); err != nil {
|
||||
return err
|
||||
}
|
||||
movedTarget = true
|
||||
} else if !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := os.Rename(temporaryPath, target); err != nil {
|
||||
if movedTarget {
|
||||
_ = os.Rename(backup, target)
|
||||
}
|
||||
return err
|
||||
}
|
||||
if movedTarget {
|
||||
if err := os.Remove(backup); err != nil && !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func transactionPaths(root string) (string, string) {
|
||||
return filepath.Join(root, transactionFileName),
|
||||
filepath.Join(root, transactionBackupFileName)
|
||||
}
|
||||
Reference in New Issue
Block a user