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) }