fix: close Pi management final review findings

This commit is contained in:
2026-08-14 20:15:21 +02:00
parent 225d13cd75
commit 900faad983
13 changed files with 621 additions and 71 deletions
+15 -5
View File
@@ -333,7 +333,11 @@ func piCommand(ctx context.Context, installation config.Installation, runner com
fmt.Fprintf(stdout, "Pi core restarted with the existing image; version %s readiness and smoke checks passed.\n", output.Sanitize(result.Version, secretValues))
return 0
case "update":
request, err := parsePiUpdateArgs(args[1:], installation.UpdateStatePath())
request, err := parsePiUpdateArgs(
args[1:],
installation.UpdateStatePath(),
installation.RestartStatePath(),
)
if err != nil {
return commandUsageError(stderr, err.Error())
}
@@ -351,7 +355,13 @@ func piCommand(ctx context.Context, installation config.Installation, runner com
if len(args) != 2 || args[1] != "--yes" {
return commandUsageError(stderr, "pi rollback requires --yes")
}
result, err := pi.Rollback(ctx, controlled, installation.UpdateStatePath(), true)
result, err := pi.Rollback(
ctx,
controlled,
installation.UpdateStatePath(),
installation.RestartStatePath(),
true,
)
if err != nil {
return piFailure(stderr, err, secretValues)
}
@@ -489,8 +499,8 @@ func parsePiConfigureArgs(args []string) (pi.Defaults, error) {
return value, nil
}
func parsePiUpdateArgs(args []string, statePath string) (pi.Request, error) {
request := pi.Request{StatePath: statePath}
func parsePiUpdateArgs(args []string, statePath, restartStatePath string) (pi.Request, error) {
request := pi.Request{StatePath: statePath, RestartStatePath: restartStatePath}
for len(args) > 0 {
switch args[0] {
case "--version":
@@ -566,7 +576,7 @@ func parsePiRestartArgs(args []string, restartStatePath, updateStatePath string)
func piFailure(stderr io.Writer, err error, secretValues []string) int {
code := 1
if errors.Is(err, pi.ErrConfirmationRequired) || errors.Is(err, pi.ErrInvalidRequest) || errors.Is(err, pi.ErrActiveSessions) || errors.Is(err, pi.ErrInterruptedUpdate) {
if errors.Is(err, pi.ErrConfirmationRequired) || errors.Is(err, pi.ErrInvalidRequest) || errors.Is(err, pi.ErrActiveSessions) || errors.Is(err, pi.ErrInterruptedUpdate) || errors.Is(err, pi.ErrInterruptedRestart) {
code = 2
}
var childExit interface{ ExitCode() int }
+53
View File
@@ -274,6 +274,59 @@ func TestRunLogsRedactsAnUnlabelledDeclaredSecret(t *testing.T) {
}
}
func TestRunPiLogsRedactsNestedAuthAndBundleScalars(t *testing.T) {
fixture := newCLIFixture(t, "")
authFile := filepath.Join(fixture.root, "pi-auth.json")
bundleFile := filepath.Join(fixture.root, "thothii.secrets")
if err := os.WriteFile(authFile, []byte(`{
"providers": {
"dummy-provider": {
"auth": {
"key": "dummy-canary-pi-log-json",
"nested": {"access": "dummy-canary-pi-log-nested"}
}
}
}
}`), 0o600); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(bundleFile, []byte(
"MODEL_API_KEY=dummy-canary-pi-log-model\n"+
"DWH_PASSWORD=dummy-canary-pi-log-dwh\n",
), 0o600); err != nil {
t.Fatal(err)
}
fixture.setEnvContents(t, "PI_AUTH_FILE="+authFile+"\nTHT_SECRETS_FILE="+bundleFile+"\n")
t.Setenv(
"THOTHCTL_FAKE_LOG",
"dummy-canary-pi-log-json dummy-canary-pi-log-nested "+
"dummy-canary-pi-log-model dummy-canary-pi-log-dwh",
)
var stdout, stderr bytes.Buffer
exitCode := run(context.Background(), []string{
"--installation", fixture.installationPath, "pi", "logs",
}, &stdout, &stderr)
if exitCode != 0 {
t.Fatalf("run() exit code = %d, stderr = %q", exitCode, stderr.String())
}
for _, canary := range []string{
"dummy-canary-pi-log-json",
"dummy-canary-pi-log-nested",
"dummy-canary-pi-log-model",
"dummy-canary-pi-log-dwh",
} {
if strings.Contains(stdout.String()+stderr.String(), canary) {
t.Fatalf("pi logs exposed scalar %q: stdout=%q stderr=%q", canary, stdout.String(), stderr.String())
}
}
if strings.Count(stdout.String(), "[REDACTED]") != 4 {
t.Fatalf("pi logs = %q, want four scalar redactions", stdout.String())
}
assertInvocationContains(t, fixture.invocations(t), "logs", "--tail", "200", "core")
}
func TestRunResolvesComposeDotenvCommentsQuotesAndInterpolationForSecretFiles(t *testing.T) {
fixture := newCLIFixture(t, "")
secretDirectory := filepath.Join(fixture.root, "secret directory")
+116 -6
View File
@@ -2,15 +2,21 @@
package output
import (
"bytes"
"encoding/json"
"errors"
"io"
"regexp"
"sort"
"strings"
"github.com/aritmolab/thothii/tools/thothctl/internal/safeio"
"github.com/compose-spec/compose-go/v2/dotenv"
)
var credentialField = regexp.MustCompile(`(?im)(\b[\w.-]*(?:password|token|key)[\w.-]*\s*[:=]\s*)(?:"[^"\r\n]*"|'[^'\r\n]*'|[^\s,;]+)`)
var credentialField = regexp.MustCompile(`(?im)((?:"|')?[\w.-]*(?:password|token|key|secret|credential)[\w.-]*(?:"|')?\s*[:=]\s*)(?:"(?:\\.|[^"\\\r\n])*"|'[^'\r\n]*'|[^\s,;}]+)`)
var dotenvAssignment = regexp.MustCompile(`(?m)^\s*(?:export\s+)?[A-Za-z_][A-Za-z0-9_.-]*\s*=`)
const maxSecretFileBytes = 64 * 1024
@@ -18,6 +24,12 @@ const maxSecretSourceFiles = 32
const maxSecretSourceBytes = 256 * 1024
const maxSecretValuesPerFile = 1024
const maxSecretValues = 4096
const maxJSONSecretDepth = 32
const maxDiagnosticDetailBytes = 512
// Sanitize redacts common credential fields and every supplied secret value.
@@ -28,6 +40,9 @@ func Sanitize(text string, secretValues []string) string {
for _, value := range values {
if value != "" {
text = strings.ReplaceAll(text, value, "[REDACTED]")
if encoded, err := json.Marshal(value); err == nil {
text = strings.ReplaceAll(text, string(encoded), "[REDACTED]")
}
}
}
return text
@@ -60,7 +75,7 @@ func SecretValuesFromFiles(paths []string) ([]string, error) {
seen := make(map[string]struct{})
var totalBytes int64
for _, path := range paths {
value, size, err := readSecretFile(path)
contents, size, err := readSecretFile(path)
if err != nil {
return nil, err
}
@@ -68,21 +83,116 @@ func SecretValuesFromFiles(paths []string) ([]string, error) {
if totalBytes > maxSecretSourceBytes {
return nil, errors.New("declared secret file could not be read")
}
if value != "" {
extracted, err := extractSecretValues(contents)
if err != nil {
return nil, errors.New("declared secret file could not be read")
}
for _, value := range extracted {
if value == "" {
continue
}
if _, exists := seen[value]; exists {
continue
}
values = append(values, value)
seen[value] = struct{}{}
if len(values) > maxSecretValues {
return nil, errors.New("declared secret file could not be read")
}
}
}
return values, nil
}
func readSecretFile(path string) (string, int64, error) {
func readSecretFile(path string) ([]byte, int64, error) {
contents, err := safeio.ReadCanonicalRegular(path, maxSecretFileBytes)
if err != nil {
return "", 0, errors.New("declared secret file could not be read")
return nil, 0, errors.New("declared secret file could not be read")
}
return strings.TrimRight(string(contents), "\r\n"), int64(len(contents)), nil
return contents, int64(len(contents)), nil
}
func extractSecretValues(contents []byte) ([]string, error) {
whole := strings.TrimRight(string(contents), "\r\n")
trimmed := bytes.TrimSpace(contents)
if len(trimmed) == 0 {
return nil, nil
}
values := make([]string, 0, 8)
if whole != "" {
values = append(values, whole)
}
if trimmed[0] == '{' || trimmed[0] == '[' {
var document any
decoder := json.NewDecoder(bytes.NewReader(trimmed))
decoder.UseNumber()
if err := decoder.Decode(&document); err != nil {
return nil, err
}
var extra any
if err := decoder.Decode(&extra); !errors.Is(err, io.EOF) {
if err == nil {
return nil, errors.New("secret JSON contains multiple documents")
}
return nil, err
}
count := 0
if err := collectJSONSecretValues(document, 0, &count, &values); err != nil {
return nil, err
}
return values, nil
}
if dotenvAssignment.Match(trimmed) {
parsed, err := dotenv.Parse(bytes.NewReader(contents))
if err != nil || len(parsed) > maxSecretValuesPerFile {
return nil, errors.New("secret dotenv bundle is invalid")
}
keys := make([]string, 0, len(parsed))
for key := range parsed {
keys = append(keys, key)
}
sort.Strings(keys)
for _, key := range keys {
if parsed[key] != "" {
values = append(values, parsed[key])
}
}
}
return values, nil
}
func collectJSONSecretValues(value any, depth int, count *int, values *[]string) error {
if depth > maxJSONSecretDepth {
return errors.New("secret JSON nesting is too deep")
}
switch typed := value.(type) {
case map[string]any:
keys := make([]string, 0, len(typed))
for key := range typed {
keys = append(keys, key)
}
sort.Strings(keys)
for _, key := range keys {
if err := collectJSONSecretValues(typed[key], depth+1, count, values); err != nil {
return err
}
}
case []any:
for _, item := range typed {
if err := collectJSONSecretValues(item, depth+1, count, values); err != nil {
return err
}
}
default:
*count++
if *count > maxSecretValuesPerFile {
return errors.New("secret JSON contains too many scalar values")
}
if scalar, ok := typed.(string); ok && scalar != "" {
*values = append(*values, scalar)
}
}
return nil
}
@@ -3,6 +3,7 @@ package output
import (
"os"
"path/filepath"
"strings"
"testing"
)
@@ -34,6 +35,113 @@ func TestSanitizeRedactsSecretFileContents(t *testing.T) {
}
}
func TestSecretValuesFromFilesRedactsNestedJSONScalars(t *testing.T) {
t.Parallel()
secretFile := filepath.Join(physicalTempDir(t), "pi-auth.json")
contents := `{
"providers": {
"dummy-provider": {
"auth": {
"key": "dummy-canary-json-key",
"tokens": {
"access": "dummy-canary-json-access",
"refresh": "dummy-canary-json-refresh"
}
}
}
}
}`
if err := os.WriteFile(secretFile, []byte(contents), 0o600); err != nil {
t.Fatal(err)
}
secrets, err := SecretValuesFromFiles([]string{secretFile})
if err != nil {
t.Fatalf("SecretValuesFromFiles() error = %v", err)
}
got := Sanitize(
"unlabelled dummy-canary-json-key dummy-canary-json-access dummy-canary-json-refresh",
secrets,
)
for _, canary := range []string{
"dummy-canary-json-key",
"dummy-canary-json-access",
"dummy-canary-json-refresh",
} {
if strings.Contains(got, canary) {
t.Fatalf("Sanitize() exposed nested JSON scalar %q: %q", canary, got)
}
}
}
func TestSecretValuesFromFilesRedactsEveryDotenvBundleValue(t *testing.T) {
t.Parallel()
secretFile := filepath.Join(physicalTempDir(t), "thothii.secrets")
contents := "MODEL_API_KEY=dummy-canary-bundle-model\n" +
"DWH_PASSWORD='dummy-canary-bundle-dwh'\n" +
"SESSION_TOKEN=\"dummy-canary-bundle-session\"\n"
if err := os.WriteFile(secretFile, []byte(contents), 0o600); err != nil {
t.Fatal(err)
}
secrets, err := SecretValuesFromFiles([]string{secretFile})
if err != nil {
t.Fatalf("SecretValuesFromFiles() error = %v", err)
}
got := Sanitize(
"unlabelled dummy-canary-bundle-model dummy-canary-bundle-dwh dummy-canary-bundle-session",
secrets,
)
for _, canary := range []string{
"dummy-canary-bundle-model",
"dummy-canary-bundle-dwh",
"dummy-canary-bundle-session",
} {
if strings.Contains(got, canary) {
t.Fatalf("Sanitize() exposed dotenv bundle scalar %q: %q", canary, got)
}
}
}
func TestSanitizeRecognizesQuotedCredentialKeys(t *testing.T) {
t.Parallel()
got := Sanitize(`{"key":"dummy-canary-quoted-key","safe":"visible"}`, nil)
if strings.Contains(got, "dummy-canary-quoted-key") || !strings.Contains(got, `"safe":"visible"`) {
t.Fatalf("Sanitize() = %q, want only the quoted credential field redacted", got)
}
}
func TestSecretValuesFromFilesRejectsMalformedJSONAndExcessiveScalars(t *testing.T) {
t.Parallel()
t.Run("malformed", func(t *testing.T) {
secretFile := filepath.Join(physicalTempDir(t), "malformed-auth.json")
if err := os.WriteFile(secretFile, []byte(`{"auth":{"key":"dummy-canary-malformed"}`), 0o600); err != nil {
t.Fatal(err)
}
if _, err := SecretValuesFromFiles([]string{secretFile}); err == nil {
t.Fatal("SecretValuesFromFiles() error = nil, want malformed-JSON failure")
}
})
t.Run("scalar bound", func(t *testing.T) {
secretFile := filepath.Join(physicalTempDir(t), "many-auth-values.json")
values := make([]string, 1025)
for index := range values {
values[index] = `"dummy-canary-value"`
}
if err := os.WriteFile(secretFile, []byte("["+strings.Join(values, ",")+"]"), 0o600); err != nil {
t.Fatal(err)
}
if _, err := SecretValuesFromFiles([]string{secretFile}); err == nil {
t.Fatal("SecretValuesFromFiles() error = nil, want scalar-count failure")
}
})
}
func TestSecretValuesFromFilesRejectsOversizedFiles(t *testing.T) {
t.Parallel()
+81 -15
View File
@@ -9,7 +9,7 @@ import (
)
var (
errInterruptedRestart = errors.New("a previous Pi restart is incomplete; recover lifecycle maintenance before restarting again")
ErrInterruptedRestart = errors.New("a previous Pi restart is incomplete; recover lifecycle maintenance before another Pi lifecycle operation")
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")
@@ -70,14 +70,7 @@ func restartWithHooks(
} else if err != nil && !errors.Is(err, os.ErrNotExist) {
return RestartResult{StatePath: request.StatePath}, err
}
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) {
if err := prepareLifecycleMutation(request.StatePath, hooks.removeFile); err != nil {
return RestartResult{StatePath: request.StatePath}, err
}
clearMaintenance := true
@@ -124,8 +117,13 @@ func restartWithHooks(
return RestartResult{StatePath: request.StatePath, Version: version}, err
}
previous.ConfigurationSHA = configured.ConfigurationSHA
transaction := lifecycleTransaction(request.StatePath)
previous.Reference = lifecycleImageTag(transaction, "restart")
if err := tagImage(ctx, runner, previous.ID, previous.Reference, "restart image pin"); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, err
}
state = State{
Transaction: lifecycleTransaction(request.StatePath),
Transaction: transaction,
Phase: PhasePreflight,
Target: Target{Version: version, Source: "restart"},
Previous: previous,
@@ -133,6 +131,11 @@ func restartWithHooks(
if err := hooks.writeState(request.StatePath, state); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, err
}
overridePath := lifecycleOverridePath(request.StatePath, transaction)
if err := writeLifecycleOverride(overridePath, previous.Reference); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, err
}
lifecycle := composeOverrideRunner{Runner: runner, path: overridePath}
if running, err := activeSessions(ctx, runner); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, err
} else if running {
@@ -144,23 +147,26 @@ func restartWithHooks(
}
mutationStarted = true
clearMaintenance = false
if err := recreateCore(ctx, runner); err != nil {
if err := recreateCoreWithoutImageChanges(ctx, lifecycle); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, recoveryRequired("Pi restart core recreation failed", err)
}
if err := ensureMaintenance(ctx, runner); err != nil {
if err := ensureMaintenance(ctx, lifecycle); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, recoveryRequired("Pi restart maintenance proof failed", err)
}
state.Phase = PhaseRecreated
if err := hooks.writeState(request.StatePath, state); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, recoveryRequired("Pi restart recreation state could not be recorded", err)
}
if err := verifyRestart(ctx, runner, version, previous); err != nil {
if err := verifyRestart(ctx, lifecycle, version, previous); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, recoveryRequired("Pi restart verification failed", err)
}
state.Phase = PhaseVerified
if err := hooks.writeState(request.StatePath, state); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, recoveryRequired("Pi restart verification state could not be recorded", err)
}
if err := hooks.removeFile(overridePath); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, recoveryRequired("Pi restart image override could not be removed", err)
}
if err := hooks.removeFile(request.StatePath); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, recoveryRequired("Pi restart recovery state could not be removed", err)
}
@@ -168,6 +174,28 @@ func restartWithHooks(
return RestartResult{StatePath: request.StatePath, Version: version}, nil
}
func recreateCoreWithoutImageChanges(ctx context.Context, runner Runner) error {
result, err := runCompose(
ctx,
runner,
"up",
"--detach",
"--wait",
"--wait-timeout",
"45",
"--no-deps",
"--force-recreate",
"--no-build",
"--pull",
"never",
"core",
)
if err != nil {
return commandError("core recreation", result, err)
}
return nil
}
func verifyRestart(ctx context.Context, runner Runner, wanted string, previous Image) error {
if err := Doctor(ctx, runner); err != nil {
return err
@@ -223,14 +251,25 @@ func RecoverLifecycleMaintenance(
if err := validateRestartRecoveryState(restartState); err != nil {
return err
}
restartOverride := lifecycleOverridePath(restartStatePath, restartState.Transaction)
if restartState.MutationStarted {
if err := ensureMaintenance(ctx, runner); err != nil {
if err := tagImage(ctx, runner, restartState.Previous.ID, restartState.Previous.Reference, "restart recovery image pin"); err != nil {
return recoveryRequired("Pi restart recovery image pin could not be restored", err)
}
if err := writeLifecycleOverride(restartOverride, restartState.Previous.Reference); err != nil {
return recoveryRequired("Pi restart recovery image override could not be restored", err)
}
lifecycle := composeOverrideRunner{Runner: runner, path: restartOverride}
if err := ensureMaintenance(ctx, lifecycle); err != nil {
return recoveryRequired("Pi restart maintenance recovery failed", err)
}
if err := verifyRestart(ctx, runner, restartState.Target.Version, restartState.Previous); err != nil {
if err := verifyRestart(ctx, lifecycle, restartState.Target.Version, restartState.Previous); err != nil {
return recoveryRequired("Pi restart recovery verification failed", err)
}
}
if err := durableRemove(restartOverride); err != nil {
return recoveryRequired("Pi restart recovery image override could not be removed", err)
}
if err := durableRemove(restartStatePath); err != nil {
return recoveryRequired("Pi restart recovery state could not be removed", err)
}
@@ -292,3 +331,30 @@ func validateRestartStatePaths(restartStatePath, updateStatePath string) error {
}
return nil
}
func pairedRestartStatePath(updateStatePath string) string {
return filepath.Join(filepath.Dir(updateStatePath), "restart-state.json")
}
func prepareLifecycleMutation(restartStatePath string, removeFile func(string) error) error {
state, err := readState(restartStatePath)
if errors.Is(err, os.ErrNotExist) {
return nil
}
if err != nil {
return fmt.Errorf("restart recovery state could not be validated: %w", err)
}
if err := validateRestartRecoveryState(state); err != nil {
return err
}
if state.Phase != PhaseVerified || !state.MutationStarted {
return ErrInterruptedRestart
}
if err := removeFile(lifecycleOverridePath(restartStatePath, state.Transaction)); err != nil {
return recoveryRequired("verified restart override could not be cleaned up", err)
}
if err := removeFile(restartStatePath); err != nil {
return recoveryRequired("verified restart state could not be cleaned up", err)
}
return nil
}
+61 -3
View File
@@ -55,15 +55,51 @@ func TestRestartDrainsRecreatesOnlyCoreAndRetainsImage(t *testing.T) {
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 ")
assertCalled(t, fake.calls, "up --detach --wait --wait-timeout 45 --no-deps --force-recreate --no-build --pull never core")
assertNotCalled(t, fake.calls, "compose build --pull")
for _, call := range fake.calls {
if strings.HasPrefix(call, "pull ") {
t.Fatalf("restart invoked direct image pull: %s", call)
}
}
assertNotCalled(t, fake.calls, "frontend")
if _, err := os.Stat(result.StatePath); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("successful restart state still exists: %v", err)
}
}
func TestRestartPinsCapturedImageWhenConfiguredTagMovesBeforeRecreate(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
hooks := defaultLifecycleHooks
write := hooks.writeState
hooks.writeState = func(path string, state State) error {
if state.Phase == PhasePreflight && state.MutationStarted {
fake.tags[fake.configuredImage] = "sha256:moved-configured-tag"
fake.imageVersions["sha256:moved-configured-tag"] = "9.99.0"
}
return write(path, state)
}
restartStatePath := filepath.Join(dir, "restart-state.json")
result, err := restartWithHooks(context.Background(), fake, RestartRequest{
StatePath: restartStatePath,
UpdateStatePath: filepath.Join(dir, "update-state.json"),
Confirm: true,
}, hooks)
if err != nil {
t.Fatalf("restartWithHooks() error = %v", err)
}
if result.Version != "0.80.3" || fake.currentImage != "sha256:old" {
t.Fatalf("restart result=%+v image=%q; want captured 0.80.3 / sha256:old", result, fake.currentImage)
}
assertCalled(t, fake.calls, "image tag sha256:old thothii-core:thothctl-")
assertCalled(t, fake.calls, "pi-lifecycle-")
if matches, globErr := filepath.Glob(filepath.Join(dir, "pi-lifecycle-*.yaml")); globErr != nil || len(matches) != 0 {
t.Fatalf("successful restart overrides = %v, error = %v; want safe cleanup", matches, globErr)
}
}
func TestRestartRefusesActiveSessionsWithoutDrain(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
@@ -193,6 +229,11 @@ func TestRestartPostRecreateFailureKeepsMaintenanceAndRecoveryState(t *testing.T
if stateErr != nil || !state.MutationStarted {
t.Fatalf("restart recovery state = %+v, %v; want durable mutation state", state, stateErr)
}
overridePath := lifecycleOverridePath(statePath, state.Transaction)
selected, overrideErr := readLifecycleOverride(overridePath)
if overrideErr != nil || selected != state.Previous.Reference || fake.tags[selected] != state.Previous.ID {
t.Fatalf("restart override = %q, %v; want retained exact image %q", selected, overrideErr, state.Previous.ID)
}
}
func TestRestartMaintenanceClearFailureRestoresRecoveryState(t *testing.T) {
@@ -233,6 +274,16 @@ func TestRecoverLifecycleMaintenanceVerifiesAndClearsRestartState(t *testing.T)
Previous: previous,
MutationStarted: true,
})
restartState, err := readState(restartStatePath)
if err != nil {
t.Fatal(err)
}
restartOverride := lifecycleOverridePath(restartStatePath, restartState.Transaction)
if err := writeLifecycleOverride(restartOverride, restartState.Previous.Reference); err != nil {
t.Fatal(err)
}
delete(fake.tags, restartState.Previous.Reference)
fake.tags[fake.configuredImage] = "sha256:moved-before-recovery"
candidate := previous
candidate.Reference = "thothii-core:thothctl-recover-candidate"
writeStateForTest(t, updateStatePath, State{
@@ -254,6 +305,13 @@ func TestRecoverLifecycleMaintenanceVerifiesAndClearsRestartState(t *testing.T)
if _, err := os.Stat(restartStatePath); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("restart recovery state still exists: %v", err)
}
if _, err := os.Stat(restartOverride); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("restart recovery override still exists: %v", err)
}
if fake.tags[restartState.Previous.Reference] != restartState.Previous.ID {
t.Fatalf("restart recovery pin = %q, want %q", fake.tags[restartState.Previous.Reference], restartState.Previous.ID)
}
assertCalled(t, fake.calls, restartOverride)
if fake.maintenance {
t.Fatal("maintenance remained active after both lifecycle states were verified")
}
+28 -9
View File
@@ -37,12 +37,13 @@ const (
// Request contains only non-secret operator inputs.
type Request struct {
StatePath string
Version string
Source Source
Image string
Confirm bool
Drain bool
StatePath string
RestartStatePath string
Version string
Source Source
Image string
Confirm bool
Drain bool
}
// Result summarizes the completed, failed, or recovered transaction without command output.
@@ -72,6 +73,12 @@ func updateWithHooks(ctx context.Context, runner Runner, request Request, hooks
if request.StatePath == "" {
return Result{}, errors.New("update state path is required")
}
if request.RestartStatePath == "" {
request.RestartStatePath = pairedRestartStatePath(request.StatePath)
}
if err := validateRestartStatePaths(request.RestartStatePath, request.StatePath); err != nil {
return Result{StatePath: request.StatePath}, err
}
lock, err := acquireLock(request.StatePath)
if err != nil {
return Result{StatePath: request.StatePath}, err
@@ -96,6 +103,9 @@ func updateWithHooks(ctx context.Context, runner Runner, request Request, hooks
}
request.Image = canonical
}
if err := prepareLifecycleMutation(request.RestartStatePath, hooks.removeFile); err != nil {
return Result{StatePath: request.StatePath}, err
}
if old, err := readState(request.StatePath); err == nil && stateNeedsRecovery(old) {
return Result{StatePath: request.StatePath}, ErrInterruptedUpdate
} else if err != nil && !errors.Is(err, os.ErrNotExist) {
@@ -254,11 +264,17 @@ func failPreparation(statePath, overridePath string, state State, cause error, h
}
// Rollback restores the image recorded in durable update state. It is safe for interrupted runs.
func Rollback(ctx context.Context, runner Runner, statePath string, confirm bool) (result Result, retErr error) {
return rollbackWithHooks(ctx, runner, statePath, confirm, defaultLifecycleHooks)
func Rollback(ctx context.Context, runner Runner, statePath, restartStatePath string, confirm bool) (result Result, retErr error) {
return rollbackWithHooks(ctx, runner, statePath, restartStatePath, confirm, defaultLifecycleHooks)
}
func rollbackWithHooks(ctx context.Context, runner Runner, statePath string, confirm bool, hooks lifecycleHooks) (result Result, retErr error) {
func rollbackWithHooks(ctx context.Context, runner Runner, statePath, restartStatePath string, confirm bool, hooks lifecycleHooks) (result Result, retErr error) {
if restartStatePath == "" {
restartStatePath = pairedRestartStatePath(statePath)
}
if err := validateRestartStatePaths(restartStatePath, statePath); err != nil {
return Result{StatePath: statePath}, err
}
lock, err := acquireLock(statePath)
if err != nil {
return Result{StatePath: statePath}, err
@@ -267,6 +283,9 @@ func rollbackWithHooks(ctx context.Context, runner Runner, statePath string, con
if !confirm {
return Result{StatePath: statePath}, ErrConfirmationRequired
}
if err := prepareLifecycleMutation(restartStatePath, hooks.removeFile); err != nil {
return Result{StatePath: statePath}, err
}
maintenanceErr := ensureMaintenance(ctx, runner)
clearMaintenance := maintenanceErr == nil
defer func() {
+100 -6
View File
@@ -114,7 +114,7 @@ func TestSuccessfulUpdateAndRollbackRemainSelectedOnFreshRecreate(t *testing.T)
if fake.currentImage != "sha256:candidate" {
t.Fatalf("fresh recreate image = %q, want verified candidate", fake.currentImage)
}
if _, err := Rollback(context.Background(), fake, statePath, true); err != nil {
if _, err := Rollback(context.Background(), fake, statePath, pairedRestartStatePath(statePath), true); err != nil {
t.Fatal(err)
}
fake.currentImage = "sha256:candidate"
@@ -209,7 +209,7 @@ func TestDigestPinnedConfiguredImageIsNeverUsedAsARollbackTagTarget(t *testing.T
if _, err := Update(context.Background(), fake, Request{StatePath: statePath, Version: "0.81.0", Source: BuildSource, Confirm: true}); err != nil {
t.Fatal(err)
}
if _, err := Rollback(context.Background(), fake, statePath, true); err != nil {
if _, err := Rollback(context.Background(), fake, statePath, pairedRestartStatePath(statePath), true); err != nil {
t.Fatal(err)
}
assertNotCalled(t, fake.calls, "image tag sha256:old "+fake.configuredImage)
@@ -306,7 +306,7 @@ func TestSuccessfulLifecycleCommandsPreserveTypedRecoveryErrorFromMaintenanceCle
Confirm: true,
})
} else {
result, err = Rollback(context.Background(), fake, statePath, true)
result, err = Rollback(context.Background(), fake, statePath, pairedRestartStatePath(statePath), true)
}
if !fake.maintenance {
@@ -382,7 +382,7 @@ func TestManualRollbackSurvivesADeadCandidateCore(t *testing.T) {
fake.coreRunning = false
fake.maintenance = false
result, err := Rollback(context.Background(), fake, statePath, true)
result, err := Rollback(context.Background(), fake, statePath, pairedRestartStatePath(statePath), true)
if err != nil || result.Phase != PhaseRolledBack {
t.Fatalf("Rollback() = %+v, %v; want restored previous core", result, err)
@@ -657,7 +657,7 @@ func TestRollbackFinalStateWriteFailureKeepsMaintenanceAndOverrideForRecovery(t
hooks := defaultLifecycleHooks
hooks.writeState = func(string, State) error { return errors.New("injected rollback state write failure") }
result, err := rollbackWithHooks(context.Background(), fake, statePath, true, hooks)
result, err := rollbackWithHooks(context.Background(), fake, statePath, pairedRestartStatePath(statePath), true, hooks)
if err == nil || result.Phase != PhaseFailed {
t.Fatalf("rollbackWithHooks() = %+v, %v; want failed durable finalization", result, err)
}
@@ -783,7 +783,7 @@ func TestRollbackRestoresInterruptedOrPreviouslyRecordedState(t *testing.T) {
previous.Reference = "thothii-core:thothctl-test-previous"
fake.tags[previous.Reference] = previous.ID
writeStateForTest(t, statePath, State{Version: 1, Phase: PhaseRecreated, Previous: previous})
result, err := Rollback(context.Background(), fake, statePath, true)
result, err := Rollback(context.Background(), fake, statePath, pairedRestartStatePath(statePath), true)
if err != nil {
t.Fatalf("Rollback() error = %v", err)
}
@@ -805,6 +805,100 @@ func TestUpdateRefusesToOverwriteInterruptedRecoveryState(t *testing.T) {
assertNotCalled(t, fake.calls, "compose")
}
func TestFailedRestartBlocksUpdateAndRollbackAndMalformedStateAlsoRejects(t *testing.T) {
for _, operation := range []string{"update", "rollback"} {
for _, restartState := range []string{"failed", "malformed"} {
t.Run(operation+"_"+restartState, func(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
updateStatePath := filepath.Join(dir, "update-state.json")
restartStatePath := filepath.Join(dir, "restart-state.json")
if restartState == "failed" {
fake.fail = "health"
_, restartErr := Restart(context.Background(), fake, RestartRequest{
StatePath: restartStatePath,
UpdateStatePath: updateStatePath,
Confirm: true,
})
var recovery *RecoveryRequiredError
if !errors.As(restartErr, &recovery) {
t.Fatalf("Restart() error = %v, want failed restart recovery state", restartErr)
}
fake.fail = ""
} else if err := os.WriteFile(restartStatePath, []byte("{malformed"), 0o600); err != nil {
t.Fatal(err)
}
if operation == "rollback" {
previous := stateImageForTest(t, fake)
writeStateForTest(t, updateStatePath, State{
Transaction: "rollback-target",
Phase: PhaseRecreated,
Target: Target{Version: fake.version, Source: string(BuildSource)},
Previous: previous,
MutationStarted: true,
})
}
fake.calls = nil
var err error
if operation == "update" {
_, err = Update(context.Background(), fake, Request{
StatePath: updateStatePath,
RestartStatePath: restartStatePath,
Version: "0.81.0",
Source: BuildSource,
Confirm: true,
})
} else {
_, err = Rollback(context.Background(), fake, updateStatePath, restartStatePath, true)
}
if err == nil {
t.Fatalf("%s accepted %s restart recovery state", operation, restartState)
}
if restartState == "failed" && !errors.Is(err, ErrInterruptedRestart) {
t.Fatalf("%s error = %v, want ErrInterruptedRestart", operation, err)
}
assertNotCalled(t, fake.calls, "compose")
})
}
}
}
func TestUpdateCleansVerifiedRestartStateBeforeNormalLifecycleWork(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
updateStatePath := filepath.Join(dir, "update-state.json")
restartStatePath := filepath.Join(dir, "restart-state.json")
state := State{
Transaction: "verified-restart",
Phase: PhaseVerified,
Target: Target{Version: fake.version, Source: "restart"},
Previous: stateImageForTest(t, fake),
MutationStarted: true,
}
writeStateForTest(t, restartStatePath, state)
overridePath := lifecycleOverridePath(restartStatePath, state.Transaction)
if err := writeLifecycleOverride(overridePath, state.Previous.Reference); err != nil {
t.Fatal(err)
}
result, err := Update(context.Background(), fake, Request{
StatePath: updateStatePath,
RestartStatePath: restartStatePath,
Version: fake.version,
Source: BuildSource,
Confirm: true,
})
if err != nil || result.Phase != PhaseNoop {
t.Fatalf("Update() = %+v, %v; want normal no-op after verified restart", result, err)
}
for _, path := range []string{restartStatePath, overridePath} {
if _, err := os.Stat(path); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("terminal restart artifact %s still exists: %v", path, err)
}
}
}
func TestRunningImageCapturesServerBindAndNamedMountIdentity(t *testing.T) {
fake := newFakeRunner()
fake.mountsJSON = `[{"Type":"bind","Source":"/srv/thothii/data","Destination":"/data","RW":true},{"Type":"bind","Source":"/srv/thothii/pi","Destination":"/home/thoth/.pi","RW":true},{"Type":"volume","Name":"sessions","Source":"/var/lib/docker/volumes/sessions/_data","Destination":"/data/sessions","RW":true}]`