Fence installation transaction durability (T-613)

This commit is contained in:
ila
2026-07-18 16:28:49 +08:00
parent 84befee70f
commit 20596a7de4
18 changed files with 697 additions and 64 deletions
+99
View File
@@ -0,0 +1,99 @@
package installer
import (
"errors"
"fmt"
"io/fs"
"os"
"path/filepath"
"sort"
)
var ErrDurability = errors.New("install durability fence failed")
type durabilityFence interface {
syncFile(file *os.File) error
syncDirectory(path string) error
}
type filesystemDurability struct{}
func (filesystemDurability) syncFile(file *os.File) error {
return file.Sync()
}
func (filesystemDurability) syncDirectory(path string) error {
return syncDirectoryPath(path)
}
func defaultDurability() durabilityFence {
return filesystemDurability{}
}
func effectiveDurability(fence durabilityFence) durabilityFence {
if fence == nil {
return defaultDurability()
}
return fence
}
func syncFileWithFence(fence durabilityFence, file *os.File, description string) error {
if err := effectiveDurability(fence).syncFile(file); err != nil {
return fmt.Errorf("%w: sync %s: %v", ErrDurability, description, err)
}
return nil
}
func syncDirectoryWithFence(fence durabilityFence, path, description string) error {
if err := effectiveDurability(fence).syncDirectory(path); err != nil {
return fmt.Errorf("%w: sync %s: %v", ErrDurability, description, err)
}
return nil
}
func syncStagingTree(fence durabilityFence, root string) error {
directories := make([]string, 0)
err := filepath.WalkDir(root, func(path string, entry fs.DirEntry, err error) error {
if err != nil {
return err
}
if entry.Type()&os.ModeSymlink != 0 {
return fmt.Errorf("%w: staging tree contains a symbolic link", ErrUnsafeInstallLayout)
}
if entry.IsDir() {
directories = append(directories, path)
}
return nil
})
if err != nil {
return fmt.Errorf("%w: walk staging tree: %w", ErrDurability, err)
}
sort.Slice(directories, func(left, right int) bool {
return len(directories[left]) > len(directories[right])
})
for _, directory := range directories {
if err := syncDirectoryWithFence(fence, directory, "staging directory"); err != nil {
return err
}
}
parent := filepath.Dir(root)
if parent != root {
if err := syncDirectoryWithFence(fence, parent, "staging parent directory"); err != nil {
return err
}
}
return nil
}
func renameManagedDirectory(
layout appLayout,
source string,
target string,
fence durabilityFence,
description string,
) error {
if err := os.Rename(source, target); err != nil {
return err
}
return syncDirectoryWithFence(fence, layout.root, description)
}
+23
View File
@@ -0,0 +1,23 @@
//go:build !windows
package installer
import (
"fmt"
"os"
)
func syncDirectoryPath(path string) error {
directory, err := os.Open(path)
if err != nil {
return fmt.Errorf("open directory: %w", err)
}
if err := directory.Sync(); err != nil {
_ = directory.Close()
return fmt.Errorf("sync directory: %w", err)
}
if err := directory.Close(); err != nil {
return fmt.Errorf("close directory: %w", err)
}
return nil
}
+334
View File
@@ -0,0 +1,334 @@
package installer
import (
"errors"
"fmt"
"os"
"path/filepath"
"strings"
"testing"
)
type durabilityEvent struct {
kind string
path string
}
type recordingDurabilityFence struct {
events []durabilityEvent
fail func(durabilityEvent) error
failNextDirectory bool
}
func (fence *recordingDurabilityFence) syncFile(file *os.File) error {
return fence.record(durabilityEvent{kind: "file", path: file.Name()})
}
func (fence *recordingDurabilityFence) syncDirectory(path string) error {
event := durabilityEvent{kind: "directory", path: path}
if fence.failNextDirectory {
fence.failNextDirectory = false
fence.events = append(fence.events, event)
return errors.New("injected directory fence failure")
}
return fence.record(event)
}
func (fence *recordingDurabilityFence) record(event durabilityEvent) error {
fence.events = append(fence.events, event)
if fence.fail != nil {
return fence.fail(event)
}
return nil
}
func TestExtractorDurabilityFencesPayloadThenStagingTree(t *testing.T) {
archivePath := writeTestZIP(t, []testZIPEntry{
{name: "app.json", body: []byte(`{"entrypoint":"bin/nested/App.exe"}`)},
{name: "payload/bin/", mode: os.ModeDir | 0o755},
{name: "payload/bin/nested/", mode: os.ModeDir | 0o755},
{name: "payload/bin/nested/App.exe", body: []byte("executable"), mode: 0o755},
{name: "payload/readme.txt", body: []byte("readme")},
})
destination := filepath.Join(t.TempDir(), "staging")
fence := &recordingDurabilityFence{}
extractor := mustExtractor(t, testLimits())
extractor.durability = fence
if _, err := extractor.ExtractFile(
archivePath,
destination,
"bin/nested/App.exe",
archiveSize(t, archivePath),
); err != nil {
t.Fatalf("ExtractFile() error = %v", err)
}
want := []durabilityEvent{
{kind: "file", path: filepath.Join(destination, "bin", "nested", "App.exe")},
{kind: "file", path: filepath.Join(destination, "readme.txt")},
{kind: "directory", path: filepath.Join(destination, "bin", "nested")},
{kind: "directory", path: filepath.Join(destination, "bin")},
{kind: "directory", path: destination},
{kind: "directory", path: filepath.Dir(destination)},
}
assertDurabilityEvents(t, fence.events, want)
}
func TestExtractorDurabilityFailuresRemoveStaging(t *testing.T) {
archivePath := writeTestZIP(t, []testZIPEntry{
{name: "app.json", body: []byte(`{"entrypoint":"App.exe"}`)},
{name: "payload/App.exe", body: []byte("executable"), mode: 0o755},
})
for _, test := range []struct {
name string
fail func(durabilityEvent) error
}{
{
name: "payload sync",
fail: func(event durabilityEvent) error {
if event.kind == "file" {
return errors.New("injected payload sync failure")
}
return nil
},
},
{
name: "staging tree sync",
fail: func(event durabilityEvent) error {
if event.kind == "directory" {
return errors.New("injected staging tree sync failure")
}
return nil
},
},
} {
t.Run(test.name, func(t *testing.T) {
destination := filepath.Join(t.TempDir(), "staging")
fence := &recordingDurabilityFence{fail: test.fail}
extractor := mustExtractor(t, testLimits())
extractor.durability = fence
_, err := extractor.ExtractFile(
archivePath,
destination,
"App.exe",
archiveSize(t, archivePath),
)
if !errors.Is(err, ErrDurability) {
t.Fatalf("ExtractFile() error = %v, want %v", err, ErrDurability)
}
assertMissing(t, destination)
})
}
}
func TestWriteTransactionFailsWhenJournalRootFenceFails(t *testing.T) {
root := t.TempDir()
layout, err := inspectAppLayout(root)
if err != nil {
t.Fatalf("inspectAppLayout() error = %v", err)
}
fence := &recordingDurabilityFence{fail: func(event durabilityEvent) error {
if event.kind == "directory" {
return errors.New("injected journal root fence failure")
}
return nil
}}
err = writeTransactionWithFence(layout, newTransaction(phasePrepared, true), fence)
if !errors.Is(err, ErrDurability) {
t.Fatalf("writeTransactionWithFence() error = %v, want %v", err, ErrDurability)
}
if len(fence.events) != 2 || fence.events[0].kind != "file" ||
fence.events[1] != (durabilityEvent{kind: "directory", path: root}) {
t.Fatalf("journal fences = %#v, want temporary file then app-root directory", fence.events)
}
if _, exists, err := loadTransaction(layout); err != nil || !exists {
t.Fatalf("loadTransaction() exists=%t error=%v, want prepared journal retained for recovery", exists, err)
}
}
func TestSwitcherFencesEachPhaseBeforeAfterStep(t *testing.T) {
root := makeInstallRoot(t, "old", "new")
fence := &recordingDurabilityFence{}
switcher := NewSwitcher(func(currentPath string) error {
assertVersion(t, currentPath, "new")
return nil
})
switcher.durability = fence
switcher.afterStep = func(step switchStep) error {
if len(fence.events) == 0 {
return fmt.Errorf("%s ran without a durability fence", step)
}
last := fence.events[len(fence.events)-1]
if last != (durabilityEvent{kind: "directory", path: root}) {
return fmt.Errorf("%s ran after %#v, want app-root directory fence", step, last)
}
return nil
}
if err := switcher.Switch(root); err != nil {
t.Fatalf("Switch() error = %v", err)
}
if !hasJournalFileFence(fence.events) {
t.Fatalf("fences = %#v, want temporary journal file sync", fence.events)
}
if countDirectoryFences(fence.events, root) < 10 {
t.Fatalf("app-root directory fences = %d, want at least 10", countDirectoryFences(fence.events, root))
}
}
func TestSwitcherRenameFenceFailureLeavesRecoverablePreparedJournal(t *testing.T) {
root := makeInstallRoot(t, "old", "new")
fence := &recordingDurabilityFence{}
switcher := NewSwitcher(func(string) error { return nil })
switcher.durability = fence
switcher.afterStep = func(step switchStep) error {
if step == stepPrepared {
fence.failNextDirectory = true
}
return nil
}
err := switcher.Switch(root)
if !errors.Is(err, ErrDurability) {
t.Fatalf("Switch() error = %v, want %v", err, ErrDurability)
}
layout, layoutErr := inspectAppLayout(root)
if layoutErr != nil {
t.Fatalf("inspectAppLayout() error = %v", layoutErr)
}
record, exists, loadErr := loadTransaction(layout)
if loadErr != nil || !exists || record.Phase != phasePrepared {
t.Fatalf("transaction = %#v exists=%t error=%v, want prepared journal", record, exists, loadErr)
}
result, err := Recover(root)
if err != nil {
t.Fatalf("Recover() error = %v", err)
}
if result.Action != RecoveryRolledBack {
t.Fatalf("Recovery 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"))
assertMissing(t, filepath.Join(root, transactionFileName))
}
func TestRollbackAndRecoveryFenceFailuresRemainRecoverable(t *testing.T) {
t.Run("rollback", func(t *testing.T) {
root := makeInstallRoot(t, "old", "new")
fence := &recordingDurabilityFence{}
switcher := NewSwitcher(func(string) error { return errors.New("health failed") })
switcher.durability = fence
switcher.afterStep = func(step switchStep) error {
if step == stepRollbackRequired {
fence.failNextDirectory = true
}
return nil
}
err := switcher.Switch(root)
var rollbackErr *RollbackError
if !errors.As(err, &rollbackErr) || !errors.Is(rollbackErr.Rollback, ErrDurability) {
t.Fatalf("Switch() error = %v, want rollback durability failure", err)
}
result, err := Recover(root)
if err != nil {
t.Fatalf("Recover() error = %v", err)
}
if result.Action != RecoveryRolledBack {
t.Fatalf("Recovery action = %q, want %q", result.Action, RecoveryRolledBack)
}
assertVersion(t, filepath.Join(root, "current"), "old")
})
t.Run("recovery", func(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)
}
fence := &recordingDurabilityFence{failNextDirectory: true}
if _, err := recoverWithFence(root, fence); !errors.Is(err, ErrDurability) {
t.Fatalf("recoverWithFence() error = %v, want %v", err, ErrDurability)
}
result, err := Recover(root)
if err != nil {
t.Fatalf("Recover() error = %v", err)
}
if result.Action != RecoveryRolledBack {
t.Fatalf("Recovery action = %q, want %q", result.Action, RecoveryRolledBack)
}
assertVersion(t, filepath.Join(root, "current"), "old")
})
}
func TestCommittedCleanupFenceFailureRemainsRecoverable(t *testing.T) {
root := makeInstallRoot(t, "old", "new")
fence := &recordingDurabilityFence{}
switcher := NewSwitcher(func(string) error { return nil })
switcher.durability = fence
switcher.afterStep = func(step switchStep) error {
if step == stepCommitted {
fence.failNextDirectory = true
}
return nil
}
err := switcher.Switch(root)
if !errors.Is(err, ErrRecoveryRequired) || !errors.Is(err, ErrDurability) {
t.Fatalf("Switch() error = %v, want recovery-required durability failure", err)
}
result, err := Recover(root)
if err != nil {
t.Fatalf("Recover() error = %v", err)
}
if result.Action != RecoveryCommitted {
t.Fatalf("Recovery action = %q, want %q", result.Action, RecoveryCommitted)
}
assertVersion(t, filepath.Join(root, "current"), "new")
assertMissing(t, filepath.Join(root, "backup"))
assertMissing(t, filepath.Join(root, transactionFileName))
}
func assertDurabilityEvents(t *testing.T, got, want []durabilityEvent) {
t.Helper()
if len(got) != len(want) {
t.Fatalf("durability events = %#v, want %#v", got, want)
}
for index := range want {
if got[index] != want[index] {
t.Fatalf("durability event %d = %#v, want %#v", index, got[index], want[index])
}
}
}
func hasJournalFileFence(events []durabilityEvent) bool {
for _, event := range events {
if event.kind == "file" && strings.HasPrefix(filepath.Base(event.path), ".install-transaction-") {
return true
}
}
return false
}
func countDirectoryFences(events []durabilityEvent, path string) int {
count := 0
for _, event := range events {
if event == (durabilityEvent{kind: "directory", path: path}) {
count++
}
}
return count
}
+35
View File
@@ -0,0 +1,35 @@
//go:build windows
package installer
import (
"fmt"
"syscall"
)
func syncDirectoryPath(path string) error {
pathPointer, err := syscall.UTF16PtrFromString(path)
if err != nil {
return fmt.Errorf("encode directory path: %w", err)
}
handle, err := syscall.CreateFile(
pathPointer,
syscall.GENERIC_READ|syscall.GENERIC_WRITE,
syscall.FILE_SHARE_READ|syscall.FILE_SHARE_WRITE|syscall.FILE_SHARE_DELETE,
nil,
syscall.OPEN_EXISTING,
syscall.FILE_FLAG_BACKUP_SEMANTICS,
0,
)
if err != nil {
return fmt.Errorf("open directory handle: %w", err)
}
if err := syscall.FlushFileBuffers(handle); err != nil {
_ = syscall.CloseHandle(handle)
return fmt.Errorf("flush directory handle: %w", err)
}
if err := syscall.CloseHandle(handle); err != nil {
return fmt.Errorf("close directory handle: %w", err)
}
return nil
}
+11
View File
@@ -0,0 +1,11 @@
//go:build windows
package installer
import "testing"
func TestSyncDirectoryPath(t *testing.T) {
if err := syncDirectoryPath(t.TempDir()); err != nil {
t.Fatalf("syncDirectoryPath() error = %v", err)
}
}
+18 -6
View File
@@ -35,7 +35,8 @@ var (
// Extractor writes only payload/ contents from a pre-verified package ZIP.
type Extractor struct {
limits Limits
limits Limits
durability durabilityFence
}
type ExtractResult struct {
@@ -56,7 +57,7 @@ func NewExtractor(limits Limits) (Extractor, error) {
if err := limits.validate(); err != nil {
return Extractor{}, err
}
return Extractor{limits: limits}, nil
return Extractor{limits: limits, durability: defaultDurability()}, nil
}
// ExtractFile requires expectedPackageSize from the verified Catalog package.
@@ -88,6 +89,7 @@ func (extractor Extractor) extract(
if err := extractor.limits.validate(); err != nil {
return ExtractResult{}, err
}
fence := effectiveDurability(extractor.durability)
normalizedEntrypoint, err := normalizeEntrypoint(entrypoint)
if err != nil {
return ExtractResult{}, err
@@ -160,21 +162,21 @@ func (extractor Extractor) extract(
readLimit++
}
copied, copyErr := io.Copy(output, io.LimitReader(source, readLimit))
closeOutputErr := output.Close()
closeSourceErr := source.Close()
if copyErr != nil {
_ = output.Close()
return ExtractResult{}, fmt.Errorf("%w: read %s: %v", ErrArchiveCorrupt, entry.archivePath, copyErr)
}
if closeOutputErr != nil {
return ExtractResult{}, fmt.Errorf("close staging file: %w", closeOutputErr)
}
if closeSourceErr != nil {
_ = output.Close()
return ExtractResult{}, fmt.Errorf("%w: close %s: %v", ErrArchiveCorrupt, entry.archivePath, closeSourceErr)
}
if copied > remaining {
_ = output.Close()
return ExtractResult{}, ErrExpandedTooLarge
}
if uint64(copied) != entry.file.UncompressedSize64 {
_ = output.Close()
return ExtractResult{}, fmt.Errorf(
"%w: %s expanded to %d bytes, header declares %d",
ErrArchiveCorrupt,
@@ -183,9 +185,19 @@ func (extractor Extractor) extract(
entry.file.UncompressedSize64,
)
}
if err := syncFileWithFence(fence, output, "staging payload"); err != nil {
_ = output.Close()
return ExtractResult{}, err
}
if err := output.Close(); err != nil {
return ExtractResult{}, fmt.Errorf("close staging file: %w", err)
}
written += copied
result.Files++
}
if err := syncStagingTree(fence, destinationRoot); err != nil {
return ExtractResult{}, err
}
result.Bytes = written
result.EntrypointPath = entrypointPath
+9 -1
View File
@@ -89,6 +89,14 @@ func inspectManagedDirectory(path string) (bool, error) {
}
func removeManagedDirectory(layout appLayout, target string) error {
return removeManagedDirectoryWithFence(layout, target, defaultDurability())
}
func removeManagedDirectoryWithFence(
layout appLayout,
target string,
fence durabilityFence,
) error {
if filepath.Dir(target) != layout.root {
return fmt.Errorf("%w: refuse removal outside app root", ErrUnsafeInstallLayout)
}
@@ -106,5 +114,5 @@ func removeManagedDirectory(layout appLayout, target string) error {
if err := os.RemoveAll(target); err != nil {
return fmt.Errorf("remove managed directory %s: %w", base, err)
}
return nil
return syncDirectoryWithFence(fence, layout.root, "app root after directory removal")
}
+36 -18
View File
@@ -3,7 +3,6 @@ package installer
import (
"errors"
"fmt"
"os"
)
var ErrRecoveryInconsistent = errors.New("install recovery state is inconsistent")
@@ -24,6 +23,11 @@ type RecoveryResult struct {
// Recover resolves an interrupted transaction from journal and directory state.
func Recover(root string) (RecoveryResult, error) {
return recoverWithFence(root, defaultDurability())
}
func recoverWithFence(root string, fence durabilityFence) (RecoveryResult, error) {
fence = effectiveDurability(fence)
layout, err := inspectAppLayout(root)
if err != nil {
return RecoveryResult{}, err
@@ -57,18 +61,18 @@ func Recover(root string) (RecoveryResult, error) {
state.staging,
)
}
if err := removeManagedDirectory(layout, layout.backup); err != nil {
if err := removeManagedDirectoryWithFence(layout, layout.backup, fence); err != nil {
return RecoveryResult{}, err
}
if err := removeTransaction(layout); err != nil {
if err := removeTransactionWithFence(layout, fence); err != nil {
return RecoveryResult{}, err
}
result.Action = RecoveryCommitted
return result, nil
case phaseRollbackRequired:
return recoverRollbackRequired(layout, record, state, result)
return recoverRollbackRequired(layout, record, state, result, fence)
case phasePrepared, phaseCurrentBackedUp, phaseStagingActivated:
return recoverUncommitted(layout, record, state, result)
return recoverUncommitted(layout, record, state, result, fence)
default:
return RecoveryResult{}, fmt.Errorf("%w: phase=%q", ErrTransactionCorrupt, record.Phase)
}
@@ -79,6 +83,7 @@ func recoverUncommitted(
record transactionRecord,
state directoryState,
result RecoveryResult,
fence durabilityFence,
) (RecoveryResult, error) {
if record.HadCurrent {
if state.backup {
@@ -89,27 +94,39 @@ func recoverUncommitted(
)
}
if state.current {
if err := os.Rename(layout.current, layout.staging); err != nil {
if err := renameManagedDirectory(
layout,
layout.current,
layout.staging,
fence,
"app root after recovery staging rename",
); err != nil {
return RecoveryResult{}, fmt.Errorf("move unverified current aside: %w", err)
}
}
if err := os.Rename(layout.backup, layout.current); err != nil {
if err := renameManagedDirectory(
layout,
layout.backup,
layout.current,
fence,
"app root after recovery current restore",
); err != nil {
return RecoveryResult{}, fmt.Errorf("restore backup during recovery: %w", err)
}
if err := removeManagedDirectory(layout, layout.staging); err != nil {
if err := removeManagedDirectoryWithFence(layout, layout.staging, fence); err != nil {
return RecoveryResult{}, err
}
if err := removeTransaction(layout); err != nil {
if err := removeTransactionWithFence(layout, fence); 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 {
if err := removeManagedDirectoryWithFence(layout, layout.staging, fence); err != nil {
return RecoveryResult{}, err
}
if err := removeTransaction(layout); err != nil {
if err := removeTransactionWithFence(layout, fence); err != nil {
return RecoveryResult{}, err
}
result.Action = RecoveryAborted
@@ -127,13 +144,13 @@ func recoverUncommitted(
ErrRecoveryInconsistent,
)
}
if err := removeManagedDirectory(layout, layout.current); err != nil {
if err := removeManagedDirectoryWithFence(layout, layout.current, fence); err != nil {
return RecoveryResult{}, err
}
if err := removeManagedDirectory(layout, layout.staging); err != nil {
if err := removeManagedDirectoryWithFence(layout, layout.staging, fence); err != nil {
return RecoveryResult{}, err
}
if err := removeTransaction(layout); err != nil {
if err := removeTransactionWithFence(layout, fence); err != nil {
return RecoveryResult{}, err
}
result.Action = RecoveryAborted
@@ -145,6 +162,7 @@ func recoverRollbackRequired(
record transactionRecord,
state directoryState,
result RecoveryResult,
fence durabilityFence,
) (RecoveryResult, error) {
if record.HadCurrent && !state.backup {
if !state.current {
@@ -153,19 +171,19 @@ func recoverRollbackRequired(
ErrRecoveryInconsistent,
)
}
if err := removeManagedDirectory(layout, layout.staging); err != nil {
if err := removeManagedDirectoryWithFence(layout, layout.staging, fence); err != nil {
return RecoveryResult{}, err
}
if err := removeTransaction(layout); err != nil {
if err := removeTransactionWithFence(layout, fence); err != nil {
return RecoveryResult{}, err
}
result.Action = RecoveryRolledBack
return result, nil
}
if err := rollbackActivated(layout, record.HadCurrent); err != nil {
if err := rollbackActivated(layout, record.HadCurrent, fence); err != nil {
return RecoveryResult{}, err
}
if err := removeTransaction(layout); err != nil {
if err := removeTransactionWithFence(layout, fence); err != nil {
return RecoveryResult{}, err
}
result.Action = RecoveryRolledBack
+61 -23
View File
@@ -42,18 +42,20 @@ func (err *RollbackError) Unwrap() error {
// Switcher activates a verified staging directory and runs an injected check.
type Switcher struct {
health HealthCheck
afterStep func(switchStep) error
health HealthCheck
afterStep func(switchStep) error
durability durabilityFence
}
func NewSwitcher(health HealthCheck) *Switcher {
return &Switcher{health: health}
return &Switcher{health: health, durability: defaultDurability()}
}
func (switcher *Switcher) Switch(root string) error {
if switcher.health == nil {
return ErrHealthCheckRequired
}
fence := effectiveDurability(switcher.durability)
layout, err := inspectAppLayout(root)
if err != nil {
return err
@@ -75,7 +77,7 @@ func (switcher *Switcher) Switch(root string) error {
}
record := newTransaction(phasePrepared, state.current)
if err := writeTransaction(layout, record); err != nil {
if err := writeTransactionWithFence(layout, record, fence); err != nil {
return err
}
if err := switcher.runStep(stepPrepared); err != nil {
@@ -83,7 +85,13 @@ func (switcher *Switcher) Switch(root string) error {
}
if state.current {
if err := os.Rename(layout.current, layout.backup); err != nil {
if err := renameManagedDirectory(
layout,
layout.current,
layout.backup,
fence,
"app root after current backup rename",
); err != nil {
return fmt.Errorf("backup current directory: %w", err)
}
if err := switcher.runStep(stepCurrentRenamed); err != nil {
@@ -91,21 +99,27 @@ func (switcher *Switcher) Switch(root string) error {
}
}
record.Phase = phaseCurrentBackedUp
if err := writeTransaction(layout, record); err != nil {
if err := writeTransactionWithFence(layout, record, fence); err != nil {
return err
}
if err := switcher.runStep(stepCurrentBackedUp); err != nil {
return err
}
if err := os.Rename(layout.staging, layout.current); err != nil {
if err := renameManagedDirectory(
layout,
layout.staging,
layout.current,
fence,
"app root after staging activation rename",
); 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 {
if err := writeTransactionWithFence(layout, record, fence); err != nil {
return err
}
if err := switcher.runStep(stepStagingActivated); err != nil {
@@ -115,35 +129,35 @@ func (switcher *Switcher) Switch(root string) error {
healthErr := switcher.health(layout.current)
if healthErr != nil {
record.Phase = phaseRollbackRequired
if err := writeTransaction(layout, record); err != nil {
if err := writeTransactionWithFence(layout, record, fence); 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 {
if err := rollbackActivated(layout, record.HadCurrent, fence); err != nil {
return &RollbackError{Health: healthErr, Rollback: err}
}
if err := removeTransaction(layout); err != nil {
if err := removeTransactionWithFence(layout, fence); 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 {
if err := writeTransactionWithFence(layout, record, fence); 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 := removeManagedDirectoryWithFence(layout, layout.backup, fence); err != nil {
return fmt.Errorf("%w: cleanup committed backup: %w", ErrRecoveryRequired, err)
}
}
if err := removeTransaction(layout); err != nil {
return fmt.Errorf("%w: %v", ErrRecoveryRequired, err)
if err := removeTransactionWithFence(layout, fence); err != nil {
return fmt.Errorf("%w: %w", ErrRecoveryRequired, err)
}
return nil
}
@@ -155,7 +169,7 @@ func (switcher *Switcher) runStep(step switchStep) error {
return switcher.afterStep(step)
}
func rollbackActivated(layout appLayout, hadCurrent bool) error {
func rollbackActivated(layout appLayout, hadCurrent bool, fence durabilityFence) error {
state, err := inspectDirectories(layout)
if err != nil {
return err
@@ -168,17 +182,35 @@ func rollbackActivated(layout appLayout, hadCurrent bool) error {
if state.staging {
return fmt.Errorf("%w: current and staging both exist", ErrRollbackFailed)
}
if err := os.Rename(layout.current, layout.staging); err != nil {
if err := renameManagedDirectory(
layout,
layout.current,
layout.staging,
fence,
"app root after rollback staging rename",
); err != nil {
return fmt.Errorf("move failed current aside: %w", err)
}
}
if err := os.Rename(layout.backup, layout.current); err != nil {
if err := renameManagedDirectory(
layout,
layout.backup,
layout.current,
fence,
"app root after rollback current restore",
); err != nil {
if _, statErr := os.Stat(layout.staging); statErr == nil {
_ = os.Rename(layout.staging, layout.current)
_ = renameManagedDirectory(
layout,
layout.staging,
layout.current,
fence,
"app root after rollback restore",
)
}
return fmt.Errorf("restore previous current: %w", err)
}
return removeManagedDirectory(layout, layout.staging)
return removeManagedDirectoryWithFence(layout, layout.staging, fence)
}
if state.backup {
@@ -188,9 +220,15 @@ func rollbackActivated(layout appLayout, hadCurrent bool) error {
if state.staging {
return fmt.Errorf("%w: current and staging both exist", ErrRollbackFailed)
}
if err := os.Rename(layout.current, layout.staging); err != nil {
if err := renameManagedDirectory(
layout,
layout.current,
layout.staging,
fence,
"app root after initial rollback rename",
); err != nil {
return fmt.Errorf("move failed initial install aside: %w", err)
}
}
return removeManagedDirectory(layout, layout.staging)
return removeManagedDirectoryWithFence(layout, layout.staging, fence)
}
+40 -3
View File
@@ -66,6 +66,14 @@ func (record transactionRecord) validate() error {
}
func writeTransaction(layout appLayout, record transactionRecord) error {
return writeTransactionWithFence(layout, record, defaultDurability())
}
func writeTransactionWithFence(
layout appLayout,
record transactionRecord,
fence durabilityFence,
) error {
if err := record.validate(); err != nil {
return err
}
@@ -79,6 +87,7 @@ func writeTransaction(layout appLayout, record transactionRecord) error {
layout.transaction,
layout.transactionBackup,
data,
fence,
); err != nil {
return fmt.Errorf("write install transaction: %w", err)
}
@@ -135,11 +144,23 @@ func ensureJSONEOF(decoder *json.Decoder) error {
}
func removeTransaction(layout appLayout) error {
return removeTransactionWithFence(layout, defaultDurability())
}
func removeTransactionWithFence(layout appLayout, fence durabilityFence) error {
removed := false
if err := os.Remove(layout.transaction); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("remove install transaction: %w", err)
} else if err == nil {
removed = true
}
if err := os.Remove(layout.transactionBackup); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("remove install transaction backup: %w", err)
} else if err == nil {
removed = true
}
if removed {
return syncDirectoryWithFence(fence, layout.root, "app root after transaction removal")
}
return nil
}
@@ -149,6 +170,7 @@ func replaceFileWithBackup(
target string,
backup string,
data []byte,
fence durabilityFence,
) error {
temporary, err := os.CreateTemp(directory, ".install-transaction-*.tmp")
if err != nil {
@@ -165,7 +187,7 @@ func replaceFileWithBackup(
temporary.Close()
return err
}
if err := temporary.Sync(); err != nil {
if err := syncFileWithFence(fence, temporary, "temporary install transaction"); err != nil {
temporary.Close()
return err
}
@@ -178,12 +200,19 @@ func replaceFileWithBackup(
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) {
if err := os.Remove(backup); err != nil {
if !os.IsNotExist(err) {
return err
}
} else if err := syncDirectoryWithFence(fence, directory, "app root after stale transaction backup removal"); err != nil {
return err
}
if err := os.Rename(target, backup); err != nil {
return err
}
if err := syncDirectoryWithFence(fence, directory, "app root after transaction backup rename"); err != nil {
return err
}
movedTarget = true
} else if !os.IsNotExist(err) {
return err
@@ -191,14 +220,22 @@ func replaceFileWithBackup(
if err := os.Rename(temporaryPath, target); err != nil {
if movedTarget {
_ = os.Rename(backup, target)
if restoreErr := os.Rename(backup, target); restoreErr == nil {
_ = syncDirectoryWithFence(fence, directory, "app root after transaction restore")
}
}
return err
}
if err := syncDirectoryWithFence(fence, directory, "app root after transaction replace"); err != nil {
return err
}
if movedTarget {
if err := os.Remove(backup); err != nil && !os.IsNotExist(err) {
return err
}
if err := syncDirectoryWithFence(fence, directory, "app root after transaction backup removal"); err != nil {
return err
}
}
return nil
}