Files
soft_quay/core/updater/updater_test.go
T

490 lines
17 KiB
Go

package updater
import (
"context"
"errors"
"os"
"path/filepath"
"testing"
"time"
)
const testRequestID = "update-1234"
type testSyncer struct {
calls int
fails map[int]error
}
func (syncer *testSyncer) SyncDirectory(string) error {
syncer.calls++
return syncer.fails[syncer.calls]
}
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
}
type testHealthFunc func(context.Context, string, string, time.Duration) error
func (function testHealthFunc) WaitForHealth(ctx context.Context, target, requestID string, timeout time.Duration) error {
return function(ctx, target, requestID, timeout)
}
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.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 TestUpdateRepairsPreparedRenameSyncFailure(t *testing.T) {
request, root := testRequest(t)
syncErr := errors.New("sync root after backup rename")
syncer := &testSyncer{fails: map[int]error{2: syncErr}}
service := NewService(&testWaiter{}, &testLauncher{}, syncer, testHealth{}, Timeouts{})
err := service.Update(context.Background(), request)
if !errors.Is(err, syncErr) {
t.Fatalf("Update() error = %v, want sync error", err)
}
assertFileContents(t, filepath.Join(root, "app", ProductExecutableName), "old")
assertFileContents(t, filepath.Join(root, "staging", testRequestID, ProductExecutableName), "new")
assertNotExists(t, filepath.Join(root, transactionFileName))
}
func TestRecoverPreparedRestoresBackupAfterInterruptedRename(t *testing.T) {
request, root := testRequest(t)
layout, err := inspectRequest(request)
if err != nil {
t.Fatal(err)
}
if err := writeTransaction(layout, phasePrepared, &testSyncer{}); err != nil {
t.Fatal(err)
}
if err := os.Rename(layout.target, layout.backup); err != nil {
t.Fatal(err)
}
if err := Recover(request.TargetDir, &testSyncer{}); err != nil {
t.Fatalf("Recover() error = %v", err)
}
assertFileContents(t, filepath.Join(root, "app", ProductExecutableName), "old")
assertNotExists(t, filepath.Join(root, transactionFileName))
}
func TestRecoverPreparedRejectsContradictoryTargetAndBackup(t *testing.T) {
request, _ := testRequest(t)
layout, err := inspectRequest(request)
if err != nil {
t.Fatal(err)
}
if err := os.Mkdir(layout.backup, 0o700); err != nil {
t.Fatal(err)
}
if err := writeTransaction(layout, phasePrepared, &testSyncer{}); err != nil {
t.Fatal(err)
}
err = Recover(request.TargetDir, &testSyncer{})
if !errors.Is(err, ErrRecoveryRequired) {
t.Fatalf("Recover() error = %v, want ErrRecoveryRequired", err)
}
if _, err := os.Lstat(layout.transaction); err != nil {
t.Fatalf("prepared transaction was removed: %v", err)
}
}
func TestRecoverPreparedRetainsJournalWhenBackupRestoreFenceFails(t *testing.T) {
request, root := testRequest(t)
layout, err := inspectRequest(request)
if err != nil {
t.Fatal(err)
}
if err := writeTransaction(layout, phasePrepared, &testSyncer{}); err != nil {
t.Fatal(err)
}
if err := os.Rename(layout.target, layout.backup); err != nil {
t.Fatal(err)
}
syncErr := errors.New("sync restored target")
err = Recover(request.TargetDir, &testSyncer{fails: map[int]error{1: syncErr}})
if !errors.Is(err, ErrRecoveryRequired) || !errors.Is(err, syncErr) {
t.Fatalf("Recover() error = %v, want recovery and sync errors", err)
}
assertFileContents(t, filepath.Join(root, "app", ProductExecutableName), "old")
if _, err := os.Lstat(layout.transaction); err != nil {
t.Fatalf("transaction missing after failed restore fence: %v", err)
}
if err := Recover(request.TargetDir, &testSyncer{}); err != nil {
t.Fatalf("Recover() retry error = %v", err)
}
assertNotExists(t, layout.transaction)
}
func TestUpdatePreservesRecoveryMaterialWhenHealthRollbackIsBlocked(t *testing.T) {
request, root := testRequest(t)
blockedStaging := filepath.Join(root, stagingDirectoryName, testRequestID)
health := testHealthFunc(func(context.Context, string, string, time.Duration) error {
if err := os.Mkdir(blockedStaging, 0o700); err != nil {
return err
}
return ErrHealthTimeout
})
service := NewService(&testWaiter{}, &testLauncher{}, &testSyncer{}, health, Timeouts{})
err := service.Update(context.Background(), request)
if !errors.Is(err, ErrHealthTimeout) || !errors.Is(err, ErrRecoveryRequired) {
t.Fatalf("Update() error = %v, want health and recovery errors", err)
}
assertFileContents(t, filepath.Join(root, "app", ProductExecutableName), "new")
assertFileContents(t, filepath.Join(root, "backups", testRequestID, ProductExecutableName), "old")
if _, err := os.Lstat(filepath.Join(root, transactionFileName)); err != nil {
t.Fatalf("transaction missing after blocked rollback: %v", err)
}
if err := os.RemoveAll(blockedStaging); err != nil {
t.Fatal(err)
}
if err := Recover(request.TargetDir, &testSyncer{}); err != nil {
t.Fatalf("Recover() after unblock error = %v", err)
}
assertFileContents(t, filepath.Join(root, "app", ProductExecutableName), "old")
assertFileContents(t, filepath.Join(root, "staging", testRequestID, ProductExecutableName), "new")
}
func TestRecoverStagingActivatedRollsBackWithoutHealth(t *testing.T) {
request, layout, root := activatedLayout(t)
if err := writeTransaction(layout, phaseStagingActivated, &testSyncer{}); err != nil {
t.Fatal(err)
}
if err := Recover(request.TargetDir, &testSyncer{}); err != nil {
t.Fatalf("Recover() error = %v", err)
}
assertFileContents(t, filepath.Join(root, "app", ProductExecutableName), "old")
assertFileContents(t, filepath.Join(root, "staging", testRequestID, ProductExecutableName), "new")
}
func TestRecoverLaunchedFinalizesMatchingHealth(t *testing.T) {
request, layout, root := activatedLayout(t)
if err := writeTransaction(layout, phaseLaunched, &testSyncer{}); err != nil {
t.Fatal(err)
}
if err := AcknowledgeHealthFromExecutable(filepath.Join(root, "app", ProductExecutableName), testRequestID, &testSyncer{}); err != nil {
t.Fatal(err)
}
if err := Recover(request.TargetDir, &testSyncer{}); err != nil {
t.Fatalf("Recover() error = %v", err)
}
assertFileContents(t, filepath.Join(root, "app", ProductExecutableName), "new")
assertNotExists(t, filepath.Join(root, "backups", testRequestID))
assertNotExists(t, filepath.Join(root, transactionFileName))
}
func TestRecoverCommittedRetriesCleanupAfterSyncFailure(t *testing.T) {
request, layout, root := activatedLayout(t)
if err := writeTransaction(layout, phaseCommitted, &testSyncer{}); err != nil {
t.Fatal(err)
}
syncErr := errors.New("sync backup parent after cleanup")
err := Recover(request.TargetDir, &testSyncer{fails: map[int]error{1: syncErr}})
if !errors.Is(err, syncErr) {
t.Fatalf("Recover() error = %v, want cleanup sync error", err)
}
assertFileContents(t, filepath.Join(root, "app", ProductExecutableName), "new")
assertNotExists(t, filepath.Join(root, "backups", testRequestID))
if _, err := os.Lstat(filepath.Join(root, transactionFileName)); err != nil {
t.Fatalf("transaction missing before cleanup retry: %v", err)
}
if err := Recover(request.TargetDir, &testSyncer{}); err != nil {
t.Fatalf("Recover() retry error = %v", err)
}
assertNotExists(t, filepath.Join(root, transactionFileName))
}
func TestRecoverPreparedTransactionAfterDirectorySyncFailure(t *testing.T) {
request, _ := testRequest(t)
layout, err := inspectRequest(request)
if err != nil {
t.Fatal(err)
}
syncErr := errors.New("sync prepared transaction")
if err := writeTransaction(layout, phasePrepared, &testSyncer{fails: map[int]error{1: syncErr}}); !errors.Is(err, syncErr) {
t.Fatalf("writeTransaction() error = %v, want sync error", err)
}
if err := Recover(request.TargetDir, &testSyncer{}); err != nil {
t.Fatalf("Recover() error = %v", err)
}
assertNotExists(t, layout.transaction)
}
func TestUpdateRejectsPreparedTransactionWriteFailureWithoutMovingApp(t *testing.T) {
request, root := testRequest(t)
if err := os.Mkdir(filepath.Join(root, transactionFileName), 0o700); err != nil {
t.Fatal(err)
}
service := NewService(&testWaiter{}, &testLauncher{}, &testSyncer{}, testHealth{}, Timeouts{})
err := service.Update(context.Background(), request)
if err == nil {
t.Fatal("Update() succeeded with a non-regular transaction target")
}
assertFileContents(t, filepath.Join(root, "app", ProductExecutableName), "old")
assertFileContents(t, filepath.Join(root, "staging", testRequestID, ProductExecutableName), "new")
}
func activatedLayout(t *testing.T) (Request, updateLayout, string) {
t.Helper()
request, root := testRequest(t)
layout, err := inspectRequest(request)
if err != nil {
t.Fatal(err)
}
if err := os.Rename(layout.target, layout.backup); err != nil {
t.Fatal(err)
}
if err := os.Rename(layout.staging, layout.target); err != nil {
t.Fatal(err)
}
return request, layout, root
}
func assertFileContents(t *testing.T, path, want string) {
t.Helper()
if got := readFile(t, path); got != want {
t.Fatalf("%s = %q, want %q", path, got, want)
}
}
func assertNotExists(t *testing.T, path string) {
t.Helper()
if _, err := os.Lstat(path); !os.IsNotExist(err) {
t.Fatalf("%s exists or could not be inspected: %v", path, err)
}
}
func testRequest(t *testing.T) (Request, string) {
t.Helper()
root := t.TempDir()
for _, directory := range []string{"app", filepath.Join("staging", testRequestID), "backups", "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)
}