Implement self-update recovery (T-403)
This commit is contained in:
@@ -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