From 3a1f2ede4a5e704de811547b4c1cf34183e69756 Mon Sep 17 00:00:00 2001 From: mptyl Date: Fri, 14 Aug 2026 18:15:00 +0200 Subject: [PATCH] fix(thothctl): harden Pi restart recovery --- tools/thothctl/internal/pi/restart.go | 83 ++++++- tools/thothctl/internal/pi/restart_test.go | 253 ++++++++++++++++++++- tools/thothctl/internal/pi/update_test.go | 5 +- 3 files changed, 326 insertions(+), 15 deletions(-) diff --git a/tools/thothctl/internal/pi/restart.go b/tools/thothctl/internal/pi/restart.go index 5b4873e8..28c2f052 100644 --- a/tools/thothctl/internal/pi/restart.go +++ b/tools/thothctl/internal/pi/restart.go @@ -13,8 +13,25 @@ var ( errRestartImageDrift = errors.New("core image changed during Pi restart") errRestartConfigurationDrift = errors.New("external endpoint configuration changed during Pi restart") errRestartMountDrift = errors.New("core persistence mount contract changed during Pi restart") + errRestartConfirmation = &restartDiagnosticError{ + message: "restart requires --yes after reviewing the planned Pi core recreation", + cause: ErrConfirmationRequired, + } + errRestartActiveSessions = &restartDiagnosticError{ + message: "active sessions must be drained before restarting Pi; use --drain only after they are complete", + cause: ErrActiveSessions, + } ) +type restartDiagnosticError struct { + message string + cause error +} + +func (e *restartDiagnosticError) Error() string { return e.message } + +func (e *restartDiagnosticError) Unwrap() error { return e.cause } + type RestartRequest struct { StatePath string UpdateStatePath string @@ -46,21 +63,23 @@ func restartWithHooks( } defer lock.Release() if !request.Confirm { - return RestartResult{StatePath: request.StatePath}, ErrConfirmationRequired + return RestartResult{StatePath: request.StatePath}, errRestartConfirmation } if state, err := readState(request.UpdateStatePath); err == nil && stateNeedsRecovery(state) { return RestartResult{StatePath: request.StatePath}, ErrInterruptedUpdate } else if err != nil && !errors.Is(err, os.ErrNotExist) { return RestartResult{StatePath: request.StatePath}, err } - if state, err := readState(request.StatePath); err == nil && state.MutationStarted { - return RestartResult{StatePath: request.StatePath}, errInterruptedRestart + if state, err := readState(request.StatePath); err == nil { + if err := validateRestartRecoveryState(state); err != nil { + return RestartResult{StatePath: request.StatePath}, err + } + if state.MutationStarted { + return RestartResult{StatePath: request.StatePath}, errInterruptedRestart + } } else if err != nil && !errors.Is(err, os.ErrNotExist) { return RestartResult{StatePath: request.StatePath}, err } - if err := setMaintenance(ctx, runner, true); err != nil { - return RestartResult{StatePath: request.StatePath}, err - } clearMaintenance := true mutationStarted := false var state State @@ -82,9 +101,12 @@ func restartWithHooks( retErr = errors.Join(retErr, fmt.Errorf("maintenance admission gate could not be cleared: %w", clearErr)) } }() + if err := setMaintenance(ctx, runner, true); err != nil { + return RestartResult{StatePath: request.StatePath}, err + } if err := waitForInactiveSessions(ctx, runner, request.Drain, hooks.sleep); err != nil { - return RestartResult{StatePath: request.StatePath}, err + return RestartResult{StatePath: request.StatePath}, restartDiagnostic(err) } if err := Doctor(ctx, runner); err != nil { return RestartResult{StatePath: request.StatePath}, err @@ -114,7 +136,7 @@ func restartWithHooks( if running, err := activeSessions(ctx, runner); err != nil { return RestartResult{StatePath: request.StatePath, Version: version}, err } else if running { - return RestartResult{StatePath: request.StatePath, Version: version}, ErrActiveSessions + return RestartResult{StatePath: request.StatePath, Version: version}, errRestartActiveSessions } state.MutationStarted = true if err := hooks.writeState(request.StatePath, state); err != nil { @@ -188,7 +210,7 @@ func RecoverLifecycleMaintenance( return err } if !confirm { - return ErrConfirmationRequired + return errRestartConfirmation } lock, err := acquireLock(restartStatePath) if err != nil { @@ -198,10 +220,10 @@ func RecoverLifecycleMaintenance( restartState, restartErr := readState(restartStatePath) if restartErr == nil { + if err := validateRestartRecoveryState(restartState); err != nil { + return err + } if restartState.MutationStarted { - if restartState.Target.Source != "restart" || restartState.Target.Version == "" { - return errors.New("restart recovery state is incomplete") - } if err := ensureMaintenance(ctx, runner); err != nil { return recoveryRequired("Pi restart maintenance recovery failed", err) } @@ -218,6 +240,43 @@ func RecoverLifecycleMaintenance( return recoverMaintenanceLocked(ctx, runner, updateStatePath) } +func restartDiagnostic(err error) error { + if errors.Is(err, ErrActiveSessions) { + return errRestartActiveSessions + } + return err +} + +func validateRestartRecoveryState(state State) error { + invalid := func(reason string) error { + return fmt.Errorf("%w: restart recovery state is invalid: %s", ErrInvalidRequest, reason) + } + if state.Transaction == "" { + return invalid("transaction is missing") + } + if state.Target.Source != "restart" || state.Target.Version == "" { + return invalid("target is not a restart with a recorded version") + } + if state.Previous.ConfigurationSHA == "" { + return invalid("previous external configuration identity is missing") + } + if state.Candidate.ID != "" || state.Candidate.Reference != "" || len(state.Candidate.Mounts) != 0 || + state.Candidate.MountFingerprint != "" || state.Candidate.ConfigurationSHA != "" { + return invalid("restart state contains image candidate metadata") + } + switch state.Phase { + case PhasePreflight: + return nil + case PhaseRecreated, PhaseVerified: + if state.MutationStarted { + return nil + } + return invalid("post-recreation phase has no mutation marker") + default: + return invalid("phase is not valid for restart") + } +} + func validateRestartStatePaths(restartStatePath, updateStatePath string) error { if restartStatePath == "" { return errors.New("restart state path is required") diff --git a/tools/thothctl/internal/pi/restart_test.go b/tools/thothctl/internal/pi/restart_test.go index ad69cc61..809100c9 100644 --- a/tools/thothctl/internal/pi/restart_test.go +++ b/tools/thothctl/internal/pi/restart_test.go @@ -20,15 +20,25 @@ func TestRestartRequiresConfirmationWithoutInvokingCompose(t *testing.T) { if !errors.Is(err, ErrConfirmationRequired) { t.Fatalf("Restart() error = %v, want ErrConfirmationRequired", err) } + if got, want := err.Error(), "restart requires --yes after reviewing the planned Pi core recreation"; got != want { + t.Fatalf("Restart() error text = %q, want %q", got, want) + } assertNotCalled(t, fake.calls, "compose") } func TestRestartDrainsRecreatesOnlyCoreAndRetainsImage(t *testing.T) { fake := newFakeRunner() - fake.activeSessions = true + fake.sessionsWire = `[{"status":"open","archived":false}]` dir := t.TempDir() hooks := defaultLifecycleHooks - hooks.sleep = func(time.Duration) { fake.activeSessions = false } + sleepCalls := 0 + hooks.sleep = func(duration time.Duration) { + sleepCalls++ + if duration != time.Second { + t.Fatalf("drain sleep = %s, want %s", duration, time.Second) + } + fake.sessionsWire = `[]` + } result, err := restartWithHooks(context.Background(), fake, RestartRequest{ StatePath: filepath.Join(dir, "restart-state.json"), @@ -42,6 +52,9 @@ func TestRestartDrainsRecreatesOnlyCoreAndRetainsImage(t *testing.T) { if result.Version != fake.version { t.Fatalf("version = %q, want %q", result.Version, fake.version) } + if sleepCalls != 1 { + t.Fatalf("drain sleep calls = %d, want 1", sleepCalls) + } assertCalled(t, fake.calls, "up --detach --wait --wait-timeout 45 --no-deps --force-recreate core") assertNotCalled(t, fake.calls, "build --pull") assertNotCalled(t, fake.calls, "pull ") @@ -64,6 +77,9 @@ func TestRestartRefusesActiveSessionsWithoutDrain(t *testing.T) { if !errors.Is(err, ErrActiveSessions) { t.Fatalf("Restart() error = %v, want ErrActiveSessions", err) } + if got, want := err.Error(), "active sessions must be drained before restarting Pi; use --drain only after they are complete"; got != want { + t.Fatalf("Restart() error text = %q, want %q", got, want) + } if fake.maintenance { t.Fatal("maintenance remained active after refusing pre-mutation restart") } @@ -72,6 +88,37 @@ func TestRestartRefusesActiveSessionsWithoutDrain(t *testing.T) { } } +func TestRestartActivationFailureClearsPreMutationMaintenance(t *testing.T) { + for _, failure := range []string{ + "maintenance-activate-durability", + "maintenance-activate-durability-without-status-flag", + } { + t.Run(failure, func(t *testing.T) { + dir := t.TempDir() + fake := newFakeRunner() + fake.fail = failure + + _, err := Restart(context.Background(), fake, RestartRequest{ + StatePath: filepath.Join(dir, "restart-state.json"), + UpdateStatePath: filepath.Join(dir, "update-state.json"), + Confirm: true, + }) + var recovery *RecoveryRequiredError + if !errors.As(err, &recovery) { + t.Fatalf("Restart() error = %v, want RecoveryRequiredError", err) + } + if fake.maintenance { + t.Fatal("pre-mutation activation failure left maintenance active") + } + if fake.recreated { + t.Fatal("pre-mutation activation failure recreated core") + } + assertCalled(t, fake.calls, "/internal/maintenance/activate") + assertCalled(t, fake.calls, "/internal/maintenance/deactivate") + }) + } +} + func TestRestartRefusesInterruptedUpdateOrRestartState(t *testing.T) { for _, stateFile := range []string{"update-state.json", "restart-state.json"} { t.Run(stateFile, func(t *testing.T) { @@ -233,6 +280,208 @@ func TestRecoverLifecycleMaintenanceVerifiesAndClearsRestartState(t *testing.T) } } +func TestRecoverLifecycleMaintenanceRejectsMalformedRestartState(t *testing.T) { + for _, test := range []struct { + name string + phase Phase + source string + }{ + {name: "recreated_without_mutation_marker", phase: PhaseRecreated, source: "restart"}, + {name: "non_restart_source", phase: PhasePreflight, source: string(BuildSource)}, + } { + t.Run(test.name, func(t *testing.T) { + dir := t.TempDir() + fake := newFakeRunner() + fake.maintenance = true + restartStatePath := filepath.Join(dir, "restart-state.json") + updateStatePath := filepath.Join(dir, "update-state.json") + writeStateForTest(t, restartStatePath, State{ + Transaction: "malformed-restart", + Phase: test.phase, + Target: Target{Version: fake.version, Source: test.source}, + Previous: stateImageForTest(t, fake), + }) + fake.calls = nil + + err := RecoverLifecycleMaintenance( + context.Background(), fake, updateStatePath, restartStatePath, true, + ) + if !errors.Is(err, ErrInvalidRequest) || !strings.Contains(err.Error(), "restart recovery state is invalid") { + t.Fatalf("RecoverLifecycleMaintenance() error = %v, want invalid restart recovery state", err) + } + if !fake.maintenance { + t.Fatal("malformed restart state reopened admission") + } + if _, stateErr := os.Stat(restartStatePath); stateErr != nil { + t.Fatalf("malformed restart state was removed: %v", stateErr) + } + assertNotCalled(t, fake.calls, "/internal/maintenance/deactivate") + }) + } +} + +func TestRestartRefusesMalformedRestartStateWithoutInvokingCompose(t *testing.T) { + dir := t.TempDir() + fake := newFakeRunner() + restartStatePath := filepath.Join(dir, "restart-state.json") + writeStateForTest(t, restartStatePath, State{ + Transaction: "malformed-restart", + Phase: PhasePreflight, + Target: Target{Version: fake.version, Source: string(BuildSource)}, + Previous: stateImageForTest(t, fake), + }) + fake.calls = nil + + _, err := Restart(context.Background(), fake, RestartRequest{ + StatePath: restartStatePath, + UpdateStatePath: filepath.Join(dir, "update-state.json"), + Confirm: true, + }) + if !errors.Is(err, ErrInvalidRequest) || !strings.Contains(err.Error(), "restart recovery state is invalid") { + t.Fatalf("Restart() error = %v, want invalid restart recovery state", err) + } + if _, stateErr := os.Stat(restartStatePath); stateErr != nil { + t.Fatalf("malformed restart state was removed: %v", stateErr) + } + assertNotCalled(t, fake.calls, "compose") +} + +func TestRestartDurabilityFailureBoundaries(t *testing.T) { + injected := errors.New("injected restart durability failure") + for _, test := range []struct { + name string + configure func(*fakeRunner, *lifecycleHooks) + wantPhase Phase + wantMutation bool + wantRecreated bool + wantMaintenance bool + wantRecovery bool + wantInjected bool + }{ + { + name: "mutation_marker_write", + configure: func(_ *fakeRunner, hooks *lifecycleHooks) { + write := hooks.writeState + hooks.writeState = func(path string, state State) error { + if state.Phase == PhasePreflight && state.MutationStarted { + return injected + } + return write(path, state) + } + }, + wantPhase: PhasePreflight, + wantInjected: true, + }, + { + name: "core_recreation", + configure: func(fake *fakeRunner, _ *lifecycleHooks) { + fake.fail = "recreate" + }, + wantPhase: PhasePreflight, + wantMutation: true, + wantRecreated: true, + wantMaintenance: true, + wantRecovery: true, + }, + { + name: "maintenance_proof", + configure: func(fake *fakeRunner, _ *lifecycleHooks) { + fake.fail = "maintenance-proof" + }, + wantPhase: PhasePreflight, + wantMutation: true, + wantRecreated: true, + wantMaintenance: true, + wantRecovery: true, + }, + { + name: "recreated_phase_write", + configure: func(_ *fakeRunner, hooks *lifecycleHooks) { + write := hooks.writeState + hooks.writeState = func(path string, state State) error { + if state.Phase == PhaseRecreated { + return injected + } + return write(path, state) + } + }, + wantPhase: PhasePreflight, + wantMutation: true, + wantRecreated: true, + wantMaintenance: true, + wantRecovery: true, + wantInjected: true, + }, + { + name: "verified_phase_write", + configure: func(_ *fakeRunner, hooks *lifecycleHooks) { + write := hooks.writeState + hooks.writeState = func(path string, state State) error { + if state.Phase == PhaseVerified { + return injected + } + return write(path, state) + } + }, + wantPhase: PhaseRecreated, + wantMutation: true, + wantRecreated: true, + wantMaintenance: true, + wantRecovery: true, + wantInjected: true, + }, + { + name: "restart_state_removal", + configure: func(_ *fakeRunner, hooks *lifecycleHooks) { + hooks.removeFile = func(string) error { return injected } + }, + wantPhase: PhaseVerified, + wantMutation: true, + wantRecreated: true, + wantMaintenance: true, + wantRecovery: true, + wantInjected: true, + }, + } { + t.Run(test.name, func(t *testing.T) { + dir := t.TempDir() + fake := newFakeRunner() + hooks := defaultLifecycleHooks + test.configure(fake, &hooks) + statePath := filepath.Join(dir, "restart-state.json") + + _, err := restartWithHooks(context.Background(), fake, RestartRequest{ + StatePath: statePath, + UpdateStatePath: filepath.Join(dir, "update-state.json"), + Confirm: true, + }, hooks) + if err == nil { + t.Fatal("restartWithHooks() error = nil, want injected boundary failure") + } + var recovery *RecoveryRequiredError + if got := errors.As(err, &recovery); got != test.wantRecovery { + t.Fatalf("recovery-required = %t, want %t; error = %v", got, test.wantRecovery, err) + } + if test.wantInjected && !errors.Is(err, injected) { + t.Fatalf("restartWithHooks() error = %v, want injected durability cause", err) + } + state, stateErr := readState(statePath) + if stateErr != nil { + t.Fatalf("readState() error = %v", stateErr) + } + if state.Phase != test.wantPhase || state.MutationStarted != test.wantMutation { + t.Fatalf("state = %+v, want phase=%q mutation=%t", state, test.wantPhase, test.wantMutation) + } + if fake.recreated != test.wantRecreated { + t.Fatalf("recreated = %t, want %t", fake.recreated, test.wantRecreated) + } + if fake.maintenance != test.wantMaintenance { + t.Fatalf("maintenance = %t, want %t", fake.maintenance, test.wantMaintenance) + } + }) + } +} + func TestRestartRejectsImageConfigurationAndMountDrift(t *testing.T) { for _, test := range []struct { failure string diff --git a/tools/thothctl/internal/pi/update_test.go b/tools/thothctl/internal/pi/update_test.go index 5b3bb851..389b459c 100644 --- a/tools/thothctl/internal/pi/update_test.go +++ b/tools/thothctl/internal/pi/update_test.go @@ -1001,6 +1001,9 @@ func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose } return compose.Result{Stdout: `{"active":false,"admissions":0}`}, nil case strings.Contains(call, "/internal/maintenance/status"): + if f.fail == "maintenance-proof" && f.recreated { + return compose.Result{Stdout: fmt.Sprintf(`{"active":%t,"admissions":0,"recoveryRequired":true}`, f.maintenance)}, nil + } if f.fail == "maintenance-activate-durability" { return compose.Result{Stdout: fmt.Sprintf(`{"active":%t,"admissions":0,"recoveryRequired":true}`, f.maintenance)}, nil } @@ -1078,7 +1081,7 @@ func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose f.maintenance = false f.dropMaintenanceAfterCandidate = false } - if f.fail == "recreate" && f.currentImage == "sha256:candidate" { + if f.fail == "recreate" && (f.built || f.recreated) { return compose.Result{ExitCode: 54}, errors.New("recreate failure") } if f.fail == "compensation" && f.rollbackPrepared {