Implement self-update recovery (T-403)

This commit is contained in:
ila
2026-07-19 22:00:38 +08:00
parent 2900deba4f
commit 13004218cf
35 changed files with 1754 additions and 16 deletions
+123
View File
@@ -0,0 +1,123 @@
package updater
import (
"fmt"
"io/fs"
"os"
"path/filepath"
)
func replaceRegularFile(directory, target, pattern string, data []byte, syncer DirectorySyncer) error {
if info, err := os.Lstat(target); err == nil {
if info.Mode()&os.ModeSymlink != 0 || !info.Mode().IsRegular() {
return fmt.Errorf("%w: file target is not a regular file", ErrUnsafeLayout)
}
} else if !os.IsNotExist(err) {
return err
}
temporary, err := os.CreateTemp(directory, pattern)
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
}
if err := os.Rename(temporaryPath, target); err != nil {
return err
}
return syncer.SyncDirectory(directory)
}
func readRegularFile(path string) ([]byte, error) {
if err := validateRegularFile(path); err != nil {
return nil, err
}
return os.ReadFile(path)
}
func removeRegularFile(path, directory string, syncer DirectorySyncer) error {
info, err := os.Lstat(path)
if os.IsNotExist(err) {
return nil
}
if err != nil {
return err
}
if info.Mode()&os.ModeSymlink != 0 || !info.Mode().IsRegular() {
return fmt.Errorf("%w: refuse to remove non-regular file", ErrUnsafeLayout)
}
if err := os.Remove(path); err != nil {
return err
}
return syncer.SyncDirectory(directory)
}
func renameDirectory(source, target string, layout updateLayout, syncer DirectorySyncer, description string) error {
if source != layout.target && source != layout.staging && source != layout.backup {
return fmt.Errorf("%w: source outside managed root", ErrUnsafeLayout)
}
if target != layout.target && target != layout.staging && target != layout.backup {
return fmt.Errorf("%w: target outside managed root", ErrUnsafeLayout)
}
if err := os.Rename(source, target); err != nil {
return fmt.Errorf("%s: %w", description, err)
}
if err := syncer.SyncDirectory(filepath.Dir(source)); err != nil {
return err
}
if filepath.Dir(target) != filepath.Dir(source) {
if err := syncer.SyncDirectory(filepath.Dir(target)); err != nil {
return err
}
}
return nil
}
func removeManagedTree(path string, layout updateLayout, syncer DirectorySyncer) error {
if filepath.Dir(path) != filepath.Join(layout.root, backupsDirectoryName) && filepath.Dir(path) != filepath.Join(layout.root, stagingDirectoryName) {
return fmt.Errorf("%w: refuse removal outside managed staging or backups", ErrUnsafeLayout)
}
if err := verifyTreeNoLinks(path); err != nil {
return err
}
if err := os.RemoveAll(path); err != nil {
return err
}
return syncer.SyncDirectory(filepath.Dir(path))
}
func verifyTreeNoLinks(path string) error {
info, err := os.Lstat(path)
if os.IsNotExist(err) {
return nil
}
if err != nil {
return err
}
if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() {
return fmt.Errorf("%w: managed tree is not a real directory", ErrUnsafeLayout)
}
return filepath.WalkDir(path, func(current string, entry fs.DirEntry, walkErr error) error {
if walkErr != nil {
return walkErr
}
if entry.Type()&os.ModeSymlink != 0 {
return fmt.Errorf("%w: managed tree contains a symbolic link", ErrUnsafeLayout)
}
return nil
})
}
+125
View File
@@ -0,0 +1,125 @@
package updater
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"os"
"path/filepath"
"time"
)
type healthRecord struct {
SchemaVersion int `json:"schema_version"`
RequestID string `json:"request_id"`
}
// FileHealthWaiter reads the fixed root health file without accepting a
// caller-provided locator.
type FileHealthWaiter struct {
PollInterval time.Duration
}
// WaitForHealth waits for a matching acknowledgement below target's root.
func (waiter FileHealthWaiter) WaitForHealth(ctx context.Context, targetDir, requestID string, timeout time.Duration) error {
layout, err := inspectTarget(targetDir)
if err != nil {
return err
}
if !validRequestID(requestID) {
return fmt.Errorf("%w: unsafe request ID", ErrInvalidRequest)
}
interval := waiter.PollInterval
if interval <= 0 {
interval = 250 * time.Millisecond
}
timer := time.NewTimer(timeout)
defer timer.Stop()
ticker := time.NewTicker(interval)
defer ticker.Stop()
for {
record, readErr := readHealth(layout.health)
if readErr == nil {
if record.RequestID != requestID {
return fmt.Errorf("%w: request ID does not match", ErrHealthInvalid)
}
return nil
}
if !os.IsNotExist(readErr) {
return readErr
}
select {
case <-ctx.Done():
return ctx.Err()
case <-timer.C:
return ErrHealthTimeout
case <-ticker.C:
}
}
}
// AcknowledgeHealthFromExecutable writes the acknowledgement only when the
// running executable is exactly <root>/app/SoftBox.exe.
func AcknowledgeHealthFromExecutable(executablePath, requestID string, syncer DirectorySyncer) error {
if syncer == nil {
return fmt.Errorf("%w: directory syncer is required", ErrInvalidRequest)
}
if !validRequestID(requestID) {
return fmt.Errorf("%w: unsafe request ID", ErrInvalidRequest)
}
if !filepath.IsAbs(executablePath) || filepath.Base(filepath.Clean(executablePath)) != ProductExecutableName {
return fmt.Errorf("%w: executable is not SoftBox.exe", ErrUnsafeLayout)
}
appDir := filepath.Dir(filepath.Clean(executablePath))
if filepath.Base(appDir) != appDirectoryName {
return fmt.Errorf("%w: executable is outside root/app", ErrUnsafeLayout)
}
layout, err := inspectTarget(appDir)
if err != nil {
return err
}
if err := requireRealDirectory(layout.target, "target"); err != nil {
return err
}
if err := validateRegularFile(executablePath); err != nil {
return fmt.Errorf("%w: %v", ErrUnsafeLayout, err)
}
transaction, exists, err := loadTransaction(layout)
if err != nil {
return err
}
if !exists || transaction.RequestID != requestID ||
(transaction.Phase != phaseStagingActivated && transaction.Phase != phaseLaunched) {
return fmt.Errorf("%w: no matching activated transaction", ErrHealthInvalid)
}
record := healthRecord{SchemaVersion: transactionSchemaVersion, RequestID: requestID}
data, err := json.Marshal(record)
if err != nil {
return err
}
data = append(data, '\n')
return replaceRegularFile(layout.root, layout.health, ".self-update-health-*.tmp", data, syncer)
}
func readHealth(path string) (healthRecord, error) {
data, err := readRegularFile(path)
if err != nil {
return healthRecord{}, err
}
decoder := json.NewDecoder(bytes.NewReader(data))
decoder.DisallowUnknownFields()
var record healthRecord
if err := decoder.Decode(&record); err != nil {
return healthRecord{}, fmt.Errorf("%w: %v", ErrHealthInvalid, err)
}
var extra interface{}
if err := decoder.Decode(&extra); err != io.EOF {
return healthRecord{}, ErrHealthInvalid
}
if record.SchemaVersion != transactionSchemaVersion || !validRequestID(record.RequestID) {
return healthRecord{}, ErrHealthInvalid
}
return record, nil
}
+184
View File
@@ -0,0 +1,184 @@
package updater
import (
"fmt"
"os"
"path/filepath"
"strings"
)
const (
appDirectoryName = "app"
stagingDirectoryName = "staging"
backupsDirectoryName = "backups"
transactionFileName = "self-update-transaction.json"
healthFileName = "self-update-health.json"
transactionSchemaVersion = 1
)
type updateLayout struct {
root string
target string
staging string
backup string
transaction string
health string
requestID string
}
func inspectRequest(request Request) (updateLayout, error) {
if request.ParentPID <= 0 {
return updateLayout{}, fmt.Errorf("%w: parent PID must be positive", ErrInvalidRequest)
}
if !validRequestID(request.RequestID) {
return updateLayout{}, fmt.Errorf("%w: unsafe request ID", ErrInvalidRequest)
}
layout, err := inspectTarget(request.TargetDir)
if err != nil {
return updateLayout{}, err
}
if err := requireRealDirectory(layout.target, "target"); err != nil {
return updateLayout{}, err
}
layout.requestID = request.RequestID
expectedStaging := filepath.Join(layout.root, stagingDirectoryName, request.RequestID)
if !filepath.IsAbs(request.StagingDir) || !samePath(request.StagingDir, expectedStaging) {
return updateLayout{}, fmt.Errorf("%w: staging path does not match fixed layout", ErrUnsafeLayout)
}
layout.staging = expectedStaging
layout.backup = filepath.Join(layout.root, backupsDirectoryName, request.RequestID)
if err := requireRealDirectory(filepath.Join(layout.root, stagingDirectoryName), "staging parent"); err != nil {
return updateLayout{}, err
}
if err := requireRealDirectory(layout.staging, "staging"); err != nil {
if os.IsNotExist(unwrapPathError(err)) {
return updateLayout{}, ErrStagingMissing
}
return updateLayout{}, err
}
return layout, nil
}
func inspectTarget(targetDir string) (updateLayout, error) {
if !filepath.IsAbs(targetDir) {
return updateLayout{}, fmt.Errorf("%w: target path must be absolute", ErrUnsafeLayout)
}
target := filepath.Clean(targetDir)
if filepath.Base(target) != appDirectoryName {
return updateLayout{}, fmt.Errorf("%w: target must be root/app", ErrUnsafeLayout)
}
root := filepath.Dir(target)
if filepath.Base(root) == "." || root == target {
return updateLayout{}, fmt.Errorf("%w: target root is invalid", ErrUnsafeLayout)
}
if err := requireRealDirectory(root, "root"); err != nil {
return updateLayout{}, err
}
if exists, err := realDirectoryState(target, "target"); err != nil {
return updateLayout{}, err
} else if !exists {
// target can be absent while recovery restores a backed-up app.
return updateLayout{root: root, target: target, transaction: filepath.Join(root, transactionFileName), health: filepath.Join(root, healthFileName)}, nil
}
return updateLayout{root: root, target: target, transaction: filepath.Join(root, transactionFileName), health: filepath.Join(root, healthFileName)}, nil
}
func (layout updateLayout) validateReady(syncer DirectorySyncer) error {
if err := requireRealDirectory(layout.target, "target"); err != nil {
return err
}
if err := requireRealDirectory(layout.staging, "staging"); err != nil {
return err
}
if err := verifyTreeNoLinks(layout.staging); err != nil {
return err
}
if err := validateRegularFile(filepath.Join(layout.staging, ProductExecutableName)); err != nil {
return fmt.Errorf("%w: staged %s: %v", ErrUnsafeLayout, ProductExecutableName, err)
}
backups := filepath.Dir(layout.backup)
if info, err := os.Lstat(backups); os.IsNotExist(err) {
if err := os.Mkdir(backups, 0o700); err != nil {
return fmt.Errorf("create backups directory: %w", err)
}
if err := syncer.SyncDirectory(layout.root); err != nil {
return fmt.Errorf("sync root after creating backups directory: %w", err)
}
} else if err != nil {
return fmt.Errorf("%w: inspect backups parent: %v", ErrUnsafeLayout, err)
} else if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() {
return fmt.Errorf("%w: backups parent is not a real directory", ErrUnsafeLayout)
}
if exists, err := realDirectoryState(layout.backup, "backup"); err != nil {
return err
} else if exists {
return ErrBackupExists
}
return nil
}
func (layout updateLayout) entrypoint() string {
return filepath.Join(layout.target, ProductExecutableName)
}
func requireRealDirectory(path string, description string) error {
exists, err := realDirectoryState(path, description)
if err != nil {
return err
}
if !exists {
return &os.PathError{Op: "lstat", Path: path, Err: os.ErrNotExist}
}
return nil
}
func realDirectoryState(path string, description 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", ErrUnsafeLayout, description, err)
}
if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() {
return false, fmt.Errorf("%w: %s is not a real directory", ErrUnsafeLayout, description)
}
return true, nil
}
func validateRegularFile(path string) error {
info, err := os.Lstat(path)
if err != nil {
return err
}
if info.Mode()&os.ModeSymlink != 0 || !info.Mode().IsRegular() {
return fmt.Errorf("file is not a regular non-symlink file")
}
return nil
}
func validRequestID(value string) bool {
if len(value) < 8 || len(value) > 64 || strings.HasPrefix(value, "-") || strings.HasSuffix(value, "-") {
return false
}
for _, char := range value {
if (char < 'a' || char > 'z') && (char < '0' || char > '9') && char != '-' {
return false
}
}
return true
}
func samePath(left string, right string) bool {
return filepath.Clean(left) == filepath.Clean(right)
}
func unwrapPathError(err error) error {
for {
pathErr, ok := err.(*os.PathError)
if !ok {
return err
}
err = pathErr.Err
}
}
+102
View File
@@ -0,0 +1,102 @@
package updater
import (
"fmt"
"path/filepath"
)
func recoverLayout(layout updateLayout, syncer DirectorySyncer) error {
record, exists, err := loadTransaction(layout)
if err != nil || !exists {
return err
}
layout.requestID = record.RequestID
layout.staging = filepath.Join(layout.root, stagingDirectoryName, record.RequestID)
layout.backup = filepath.Join(layout.root, backupsDirectoryName, record.RequestID)
switch record.Phase {
case phasePrepared:
return removeTransaction(layout, syncer)
case phaseTargetBackedUp:
if err := restoreBackup(layout, syncer); err != nil {
return err
}
return removeTransaction(layout, syncer)
case phaseStagingActivated, phaseLaunched:
if healthMatches(layout, record.RequestID) {
if err := removeManagedTree(layout.backup, layout, syncer); err != nil {
return fmt.Errorf("finalize recovered self-update: %w", err)
}
return removeTransaction(layout, syncer)
}
return rollbackLayout(layout, syncer)
case phaseCommitted:
if err := removeManagedTree(layout.backup, layout, syncer); err != nil {
return fmt.Errorf("finalize committed self-update: %w", err)
}
return removeTransaction(layout, syncer)
default:
return ErrTransactionCorrupt
}
}
func rollbackLayout(layout updateLayout, syncer DirectorySyncer) error {
backupExists, err := realDirectoryState(layout.backup, "backup")
if err != nil || !backupExists {
return fmt.Errorf("%w: old app backup is unavailable", ErrRecoveryRequired)
}
targetExists, err := realDirectoryState(layout.target, "target")
if err != nil {
return err
}
if targetExists {
stagingExists, err := realDirectoryState(layout.staging, "staging")
if err != nil {
return err
}
if stagingExists {
return fmt.Errorf("%w: staged and active app both exist", ErrRecoveryRequired)
}
if err := renameDirectory(layout.target, layout.staging, layout, syncer, "preserve failed staged app"); err != nil {
return fmt.Errorf("%w: %v", ErrRecoveryRequired, err)
}
}
if err := renameDirectory(layout.backup, layout.target, layout, syncer, "restore previous app"); err != nil {
return fmt.Errorf("%w: %v", ErrRecoveryRequired, err)
}
if err := removeTransaction(layout, syncer); err != nil {
return err
}
if err := removeHealth(layout, syncer); err != nil {
return err
}
return nil
}
func restoreBackup(layout updateLayout, syncer DirectorySyncer) error {
targetExists, err := realDirectoryState(layout.target, "target")
if err != nil {
return err
}
if targetExists {
return fmt.Errorf("%w: target exists before backup restoration", ErrRecoveryRequired)
}
if exists, err := realDirectoryState(layout.backup, "backup"); err != nil || !exists {
if err != nil {
return err
}
return fmt.Errorf("%w: backup is unavailable", ErrRecoveryRequired)
}
if err := renameDirectory(layout.backup, layout.target, layout, syncer, "restore previous app"); err != nil {
return fmt.Errorf("%w: %v", ErrRecoveryRequired, err)
}
return nil
}
func healthMatches(layout updateLayout, requestID string) bool {
record, err := readHealth(layout.health)
return err == nil && record.RequestID == requestID
}
func removeHealth(layout updateLayout, syncer DirectorySyncer) error {
return removeRegularFile(layout.health, layout.root, syncer)
}
+87
View File
@@ -0,0 +1,87 @@
package updater
import (
"bytes"
"encoding/json"
"fmt"
"io"
"os"
)
type transactionPhase string
const (
phasePrepared transactionPhase = "prepared"
phaseTargetBackedUp transactionPhase = "target_backed_up"
phaseStagingActivated transactionPhase = "staging_activated"
phaseLaunched transactionPhase = "launched"
phaseCommitted transactionPhase = "committed"
)
type transaction struct {
SchemaVersion int `json:"schema_version"`
RequestID string `json:"request_id"`
Phase transactionPhase `json:"phase"`
}
func (record transaction) validate() error {
if record.SchemaVersion != transactionSchemaVersion || !validRequestID(record.RequestID) {
return ErrTransactionCorrupt
}
switch record.Phase {
case phasePrepared, phaseTargetBackedUp, phaseStagingActivated, phaseLaunched, phaseCommitted:
return nil
default:
return ErrTransactionCorrupt
}
}
func writeTransaction(layout updateLayout, phase transactionPhase, syncer DirectorySyncer) error {
record := transaction{SchemaVersion: transactionSchemaVersion, RequestID: layout.requestID, Phase: phase}
if err := record.validate(); err != nil {
return err
}
data, err := json.Marshal(record)
if err != nil {
return fmt.Errorf("encode self-update transaction: %w", err)
}
data = append(data, '\n')
if err := replaceRegularFile(layout.root, layout.transaction, ".self-update-transaction-*.tmp", data, syncer); err != nil {
return fmt.Errorf("write self-update transaction: %w", err)
}
return nil
}
func loadTransaction(layout updateLayout) (transaction, bool, error) {
data, err := readRegularFile(layout.transaction)
if os.IsNotExist(err) {
return transaction{}, false, nil
}
if err != nil {
return transaction{}, false, fmt.Errorf("read self-update transaction: %w", err)
}
decoder := json.NewDecoder(bytes.NewReader(data))
decoder.DisallowUnknownFields()
var record transaction
if err := decoder.Decode(&record); err != nil {
return transaction{}, false, fmt.Errorf("%w: %v", ErrTransactionCorrupt, err)
}
var extra interface{}
if err := decoder.Decode(&extra); err != io.EOF {
if err == nil {
return transaction{}, false, fmt.Errorf("%w: trailing JSON value", ErrTransactionCorrupt)
}
return transaction{}, false, fmt.Errorf("%w: trailing data: %v", ErrTransactionCorrupt, err)
}
if err := record.validate(); err != nil {
return transaction{}, false, err
}
return record, true, nil
}
func removeTransaction(layout updateLayout, syncer DirectorySyncer) error {
if err := removeRegularFile(layout.transaction, layout.root, syncer); err != nil {
return fmt.Errorf("remove self-update transaction: %w", err)
}
return nil
}
+198
View File
@@ -0,0 +1,198 @@
// Package updater performs the constrained on-disk activation of a prepared
// SoftBox self-update. It deliberately does not download, verify, or select an
// update package.
package updater
import (
"context"
"errors"
"fmt"
"time"
)
const (
// ProductExecutableName is the only executable the updater may start.
ProductExecutableName = "SoftBox.exe"
// InternalHealthFlag is accepted only by SoftBox itself after an update.
InternalHealthFlag = "--softbox-update-health"
)
var (
ErrInvalidRequest = errors.New("invalid self-update request")
ErrUnsafeLayout = errors.New("unsafe self-update layout")
ErrStagingMissing = errors.New("self-update staging directory is missing")
ErrBackupExists = errors.New("self-update backup directory already exists")
ErrTransactionCorrupt = errors.New("self-update transaction is corrupt")
ErrRecoveryRequired = errors.New("self-update recovery is required")
ErrParentWait = errors.New("wait for SoftBox parent process")
ErrLaunch = errors.New("launch updated SoftBox")
ErrHealthTimeout = errors.New("updated SoftBox health confirmation timed out")
ErrHealthInvalid = errors.New("updated SoftBox health confirmation is invalid")
)
// Request identifies one prepared self-update. StagingDir and TargetDir are
// both checked against the fixed layout; callers cannot choose arbitrary move
// endpoints.
type Request struct {
ParentPID int
StagingDir string
TargetDir string
RequestID string
}
// StartCommand contains the only process launch allowed by this package.
// Platform adapters must pass precisely the fixed internal health flag.
type StartCommand struct {
Entrypoint string
WorkingDirectory string
HealthRequestID string
}
// ProcessWaiter waits for the old main-process PID to end naturally.
type ProcessWaiter interface {
WaitForProcessExit(context.Context, int, time.Duration) error
}
// Launcher starts the verified new main executable without a shell.
type Launcher interface {
StartSelfUpdate(StartCommand) (int, error)
}
// DirectorySyncer is implemented by the platform boundary. Directory flushing
// requires a Windows handle implementation and must not leak into core.
type DirectorySyncer interface {
SyncDirectory(path string) error
}
// HealthWaiter observes the minimal acknowledgement written by the new main
// process. It must only accept the exact request ID.
type HealthWaiter interface {
WaitForHealth(context.Context, string, string, time.Duration) error
}
// Timeouts controls externally-blocking update operations.
type Timeouts struct {
ParentExit time.Duration
Health time.Duration
}
// Service owns one constrained self-update orchestration.
type Service struct {
waiter ProcessWaiter
launcher Launcher
syncer DirectorySyncer
health HealthWaiter
timeouts Timeouts
}
// NewService constructs an updater. Missing dependencies are reported by
// Update, keeping command composition straightforward.
func NewService(
waiter ProcessWaiter,
launcher Launcher,
syncer DirectorySyncer,
health HealthWaiter,
timeouts Timeouts,
) *Service {
if timeouts.ParentExit <= 0 {
timeouts.ParentExit = 2 * time.Minute
}
if timeouts.Health <= 0 {
timeouts.Health = 45 * time.Second
}
return &Service{waiter: waiter, launcher: launcher, syncer: syncer, health: health, timeouts: timeouts}
}
// Update waits for the old process, recovers a previous interrupted switch if
// needed, and activates the prepared directory. It never kills a process.
func (service *Service) Update(ctx context.Context, request Request) error {
if service == nil || service.waiter == nil || service.launcher == nil || service.syncer == nil || service.health == nil {
return fmt.Errorf("%w: updater dependencies are incomplete", ErrInvalidRequest)
}
layout, err := inspectRequest(request)
if err != nil {
return err
}
if err := service.waiter.WaitForProcessExit(ctx, request.ParentPID, service.timeouts.ParentExit); err != nil {
return fmt.Errorf("%w: %w", ErrParentWait, err)
}
if err := recoverLayout(layout, service.syncer); err != nil {
return err
}
if err := layout.validateReady(service.syncer); err != nil {
return err
}
if err := removeHealth(layout, service.syncer); err != nil {
return err
}
if err := writeTransaction(layout, phasePrepared, service.syncer); err != nil {
return err
}
if err := renameDirectory(layout.target, layout.backup, layout, service.syncer, "back up current app"); err != nil {
return service.failBeforeActivation(err)
}
if err := writeTransaction(layout, phaseTargetBackedUp, service.syncer); err != nil {
return service.rollback(layout, err)
}
if err := renameDirectory(layout.staging, layout.target, layout, service.syncer, "activate staged app"); err != nil {
return service.rollback(layout, err)
}
if err := writeTransaction(layout, phaseStagingActivated, service.syncer); err != nil {
return service.rollback(layout, err)
}
entrypoint := layout.entrypoint()
if err := validateRegularFile(entrypoint); err != nil {
return service.rollback(layout, fmt.Errorf("%w: %v", ErrUnsafeLayout, err))
}
if _, err := service.launcher.StartSelfUpdate(StartCommand{
Entrypoint: entrypoint, WorkingDirectory: layout.target, HealthRequestID: layout.requestID,
}); err != nil {
return service.rollback(layout, fmt.Errorf("%w: %w", ErrLaunch, err))
}
if err := writeTransaction(layout, phaseLaunched, service.syncer); err != nil {
return service.rollback(layout, err)
}
if err := service.health.WaitForHealth(ctx, layout.target, layout.requestID, service.timeouts.Health); err != nil {
return service.rollback(layout, err)
}
if err := writeTransaction(layout, phaseCommitted, service.syncer); err != nil {
return service.rollback(layout, err)
}
if err := removeManagedTree(layout.backup, layout, service.syncer); err != nil {
return fmt.Errorf("commit self-update: %w", err)
}
if err := removeTransaction(layout, service.syncer); err != nil {
return err
}
if err := removeHealth(layout, service.syncer); err != nil {
return err
}
return nil
}
func (service *Service) failBeforeActivation(cause error) error {
return fmt.Errorf("prepare self-update: %w", cause)
}
func (service *Service) rollback(layout updateLayout, cause error) error {
rollbackErr := rollbackLayout(layout, service.syncer)
if rollbackErr != nil {
return errors.Join(cause, rollbackErr)
}
return cause
}
// Recover restores or finalizes a previously interrupted transaction for the
// fixed target directory. Callers must wait for any old parent process first.
func Recover(targetDir string, syncer DirectorySyncer) error {
if syncer == nil {
return fmt.Errorf("%w: directory syncer is required", ErrInvalidRequest)
}
layout, err := inspectTarget(targetDir)
if err != nil {
return err
}
return recoverLayout(layout, syncer)
}
+254
View File
@@ -0,0 +1,254 @@
package updater
import (
"context"
"errors"
"os"
"path/filepath"
"testing"
"time"
)
const testRequestID = "update-1234"
type testSyncer struct{ err error }
func (syncer testSyncer) SyncDirectory(string) error { return syncer.err }
type testWaiter struct {
err error
calls int
}
func (waiter *testWaiter) WaitForProcessExit(context.Context, int, time.Duration) error {
waiter.calls++
return waiter.err
}
type testLauncher struct {
command StartCommand
err error
}
func (launcher *testLauncher) StartSelfUpdate(command StartCommand) (int, error) {
launcher.command = command
return 42, launcher.err
}
type testHealth struct{ err error }
func (health testHealth) WaitForHealth(context.Context, string, string, time.Duration) error {
return health.err
}
func TestUpdateActivatesOnlyFixedLayoutAfterHealth(t *testing.T) {
request, root := testRequest(t)
waiter := &testWaiter{}
launcher := &testLauncher{}
service := NewService(waiter, launcher, testSyncer{}, testHealth{}, Timeouts{})
if err := service.Update(context.Background(), request); err != nil {
t.Fatalf("Update() error = %v", err)
}
if waiter.calls != 1 {
t.Fatalf("wait calls = %d, want 1", waiter.calls)
}
if got := readFile(t, filepath.Join(root, "app", ProductExecutableName)); got != "new" {
t.Fatalf("activated executable = %q, want new", got)
}
if got := launcher.command.Entrypoint; got != filepath.Join(root, "app", ProductExecutableName) {
t.Fatalf("launch entrypoint = %q", got)
}
if launcher.command.WorkingDirectory != filepath.Join(root, "app") || launcher.command.HealthRequestID != testRequestID {
t.Fatalf("launch command = %#v, want fixed app directory and request ID", launcher.command)
}
for _, path := range []string{
filepath.Join(root, "backups", testRequestID),
filepath.Join(root, transactionFileName),
filepath.Join(root, healthFileName),
} {
if _, err := os.Lstat(path); !os.IsNotExist(err) {
t.Fatalf("%s remains after commit: %v", path, err)
}
}
if got := readFile(t, filepath.Join(root, "data", "keep.txt")); got != "data" {
t.Fatalf("data changed: %q", got)
}
if got := readFile(t, filepath.Join(root, "licenses", "keep.txt")); got != "licenses" {
t.Fatalf("licenses changed: %q", got)
}
}
func TestUpdateRestoresOldAppWhenHealthFails(t *testing.T) {
request, root := testRequest(t)
service := NewService(&testWaiter{}, &testLauncher{}, testSyncer{}, testHealth{err: ErrHealthTimeout}, Timeouts{})
err := service.Update(context.Background(), request)
if !errors.Is(err, ErrHealthTimeout) {
t.Fatalf("Update() error = %v, want ErrHealthTimeout", err)
}
if got := readFile(t, filepath.Join(root, "app", ProductExecutableName)); got != "old" {
t.Fatalf("restored executable = %q, want old", got)
}
if got := readFile(t, filepath.Join(root, "staging", testRequestID, ProductExecutableName)); got != "new" {
t.Fatalf("preserved staged executable = %q, want new", got)
}
if _, err := os.Lstat(filepath.Join(root, transactionFileName)); !os.IsNotExist(err) {
t.Fatalf("transaction remains after successful rollback: %v", err)
}
}
func TestUpdateDoesNotTouchLayoutWhenParentWaitFails(t *testing.T) {
request, root := testRequest(t)
waitErr := errors.New("permission denied")
service := NewService(&testWaiter{err: waitErr}, &testLauncher{}, testSyncer{}, testHealth{}, Timeouts{})
err := service.Update(context.Background(), request)
if !errors.Is(err, ErrParentWait) || !errors.Is(err, waitErr) {
t.Fatalf("Update() error = %v, want wrapped parent wait error", err)
}
if got := readFile(t, filepath.Join(root, "app", ProductExecutableName)); got != "old" {
t.Fatalf("target changed after wait failure: %q", got)
}
if got := readFile(t, filepath.Join(root, "staging", testRequestID, ProductExecutableName)); got != "new" {
t.Fatalf("staging changed after wait failure: %q", got)
}
if _, err := os.Lstat(filepath.Join(root, transactionFileName)); !os.IsNotExist(err) {
t.Fatalf("transaction created after wait failure: %v", err)
}
}
func TestUpdateRejectsCrossLayoutWithoutWaiting(t *testing.T) {
request, root := testRequest(t)
request.StagingDir = filepath.Join(root, "outside", testRequestID)
waiter := &testWaiter{}
service := NewService(waiter, &testLauncher{}, testSyncer{}, testHealth{}, Timeouts{})
err := service.Update(context.Background(), request)
if !errors.Is(err, ErrUnsafeLayout) {
t.Fatalf("Update() error = %v, want ErrUnsafeLayout", err)
}
if waiter.calls != 0 {
t.Fatalf("wait calls = %d, want 0", waiter.calls)
}
}
func TestUpdateRejectsStagingSymlinkWithoutWaiting(t *testing.T) {
request, root := testRequest(t)
staging := filepath.Join(root, "staging", testRequestID)
if err := os.RemoveAll(staging); err != nil {
t.Fatal(err)
}
outside := filepath.Join(root, "outside")
if err := os.Mkdir(outside, 0o700); err != nil {
t.Fatal(err)
}
writeFile(t, filepath.Join(outside, ProductExecutableName), "new")
if err := os.Symlink(outside, staging); err != nil {
t.Skipf("symlink unavailable: %v", err)
}
waiter := &testWaiter{}
service := NewService(waiter, &testLauncher{}, testSyncer{}, testHealth{}, Timeouts{})
err := service.Update(context.Background(), request)
if !errors.Is(err, ErrUnsafeLayout) {
t.Fatalf("Update() error = %v, want ErrUnsafeLayout", err)
}
if waiter.calls != 0 {
t.Fatalf("wait calls = %d, want 0", waiter.calls)
}
}
func TestRecoverRestoresTargetBackedUpTransaction(t *testing.T) {
request, root := testRequest(t)
layout, err := inspectRequest(request)
if err != nil {
t.Fatal(err)
}
if err := os.Mkdir(filepath.Join(root, backupsDirectoryName), 0o700); err != nil {
t.Fatal(err)
}
if err := os.Rename(layout.target, layout.backup); err != nil {
t.Fatal(err)
}
if err := writeTransaction(layout, phaseTargetBackedUp, testSyncer{}); err != nil {
t.Fatal(err)
}
if err := Recover(request.TargetDir, testSyncer{}); err != nil {
t.Fatalf("Recover() error = %v", err)
}
if got := readFile(t, filepath.Join(root, "app", ProductExecutableName)); got != "old" {
t.Fatalf("recovered executable = %q, want old", got)
}
if _, err := os.Lstat(filepath.Join(root, transactionFileName)); !os.IsNotExist(err) {
t.Fatalf("transaction remains after recovery: %v", err)
}
}
func TestAcknowledgeHealthFromExecutableUsesFixedLocator(t *testing.T) {
request, root := testRequest(t)
layout, err := inspectRequest(request)
if err != nil {
t.Fatal(err)
}
if err := writeTransaction(layout, phaseStagingActivated, testSyncer{}); err != nil {
t.Fatal(err)
}
executable := filepath.Join(root, "app", ProductExecutableName)
if err := AcknowledgeHealthFromExecutable(executable, testRequestID, testSyncer{}); err != nil {
t.Fatalf("AcknowledgeHealthFromExecutable() error = %v", err)
}
record, err := readHealth(filepath.Join(root, healthFileName))
if err != nil || record.RequestID != testRequestID {
t.Fatalf("health record = %#v, %v", record, err)
}
if err := AcknowledgeHealthFromExecutable(filepath.Join(root, "outside", ProductExecutableName), testRequestID, testSyncer{}); !errors.Is(err, ErrUnsafeLayout) {
t.Fatalf("outside acknowledgement error = %v, want ErrUnsafeLayout", err)
}
if err := AcknowledgeHealthFromExecutable(executable, "other-1234", testSyncer{}); !errors.Is(err, ErrHealthInvalid) {
t.Fatalf("unmatched acknowledgement error = %v, want ErrHealthInvalid", err)
}
}
func TestFileHealthWaiterRejectsWrongRequestID(t *testing.T) {
_, root := testRequest(t)
writeFile(t, filepath.Join(root, healthFileName), "{\"schema_version\":1,\"request_id\":\"other-1234\"}\n")
err := (FileHealthWaiter{PollInterval: time.Millisecond}).WaitForHealth(context.Background(), filepath.Join(root, "app"), testRequestID, time.Second)
if !errors.Is(err, ErrHealthInvalid) {
t.Fatalf("WaitForHealth() error = %v, want ErrHealthInvalid", err)
}
}
func testRequest(t *testing.T) (Request, string) {
t.Helper()
root := t.TempDir()
for _, directory := range []string{"app", filepath.Join("staging", testRequestID), "data", "licenses"} {
if err := os.MkdirAll(filepath.Join(root, directory), 0o700); err != nil {
t.Fatal(err)
}
}
writeFile(t, filepath.Join(root, "app", ProductExecutableName), "old")
writeFile(t, filepath.Join(root, "staging", testRequestID, ProductExecutableName), "new")
writeFile(t, filepath.Join(root, "data", "keep.txt"), "data")
writeFile(t, filepath.Join(root, "licenses", "keep.txt"), "licenses")
return Request{
ParentPID: 1, TargetDir: filepath.Join(root, "app"),
StagingDir: filepath.Join(root, "staging", testRequestID), RequestID: testRequestID,
}, root
}
func writeFile(t *testing.T, path string, content string) {
t.Helper()
if err := os.WriteFile(path, []byte(content), 0o600); err != nil {
t.Fatal(err)
}
}
func readFile(t *testing.T, path string) string {
t.Helper()
data, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
return string(data)
}