Implement self-update recovery (T-403)
This commit is contained in:
@@ -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
|
||||
})
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user