Prototype atomic install recovery (T-103)

This commit is contained in:
ila
2026-07-16 16:32:58 +08:00
parent 0ffeff63e1
commit 9da72caa01
10 changed files with 1075 additions and 11 deletions
+110
View File
@@ -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
}
+173
View File
@@ -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
}
+196
View File
@@ -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)
}
+295
View File
@@ -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)
}
}
+209
View File
@@ -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)
}