fix: harden Pi lifecycle recovery

This commit is contained in:
2026-08-04 20:28:09 +02:00
parent 5b3ce93e31
commit 5ba2821a1b
24 changed files with 1764 additions and 384 deletions
+128 -71
View File
@@ -21,6 +21,11 @@ type Defaults struct {
Thinking string `json:"thinking"`
}
type ModelOption struct {
Provider string `json:"provider"`
ID string `json:"id"`
}
var internalIdentityHeaders = []string{
"-H", "x-thoth-principal-issuer: thothctl",
"-H", "x-thoth-principal-subject: thothctl-maintenance",
@@ -30,7 +35,7 @@ var internalIdentityHeaders = []string{
// Configure changes the backend's real installation settings through a core-side helper. It
// deliberately has no secret or endpoint input: external endpoints remain Compose-owned.
func Configure(ctx context.Context, runner Runner, value Defaults) (retErr error) {
func Configure(ctx context.Context, runner Runner, value Defaults) error {
if !choicePattern.MatchString(value.Provider) || !choicePattern.MatchString(value.Model) {
return errors.New("provider and model must be supported identifiers")
}
@@ -38,24 +43,15 @@ func Configure(ctx context.Context, runner Runner, value Defaults) (retErr error
return errors.New("thinking must be low, medium, or high")
}
before, err := renderedCore(ctx, runner)
if err != nil { return err }
args := append([]string{"exec", "-T", "core", "curl", "-fsS"}, internalIdentityHeaders...)
args = append(args, "http://127.0.0.1:8787/models")
models, err := runCompose(ctx, runner, args...)
if err != nil {
return commandError("Pi options check", models, err)
return err
}
var payload struct {
Models []struct {
Provider string `json:"provider"`
ID string `json:"id"`
} `json:"models"`
}
if json.Unmarshal([]byte(models.Stdout), &payload) != nil || len(payload.Models) == 0 {
return errors.New("Pi options response is invalid or empty")
options, err := ConfigurationOptions(ctx, runner)
if err != nil {
return err
}
found := false
for _, model := range payload.Models {
for _, model := range options {
if model.Provider == value.Provider && model.ID == value.Model {
found = true
}
@@ -66,31 +62,86 @@ func Configure(ctx context.Context, runner Runner, value Defaults) (retErr error
settingsArgs := append([]string{"exec", "-T", "core", "curl", "-fsS"}, internalIdentityHeaders...)
settingsArgs = append(settingsArgs, "http://127.0.0.1:8787/settings")
oldResult, err := runCompose(ctx, runner, settingsArgs...)
if err != nil { return commandError("Pi installation settings capture", oldResult, err) }
if err != nil {
return commandError("Pi installation settings capture", oldResult, err)
}
var old Defaults
if json.Unmarshal([]byte(oldResult.Stdout), &old) != nil || old.Provider == "" || old.Model == "" || old.Thinking == "" { return errors.New("Pi installation settings capture is invalid") }
wrote := false
defer func() {
if retErr != nil && wrote {
result, restoreErr := runCompose(context.Background(), runner, "exec", "-T", "core", "node", "/app/backend/dist/settings/settings-cli.js", "--provider", old.Provider, "--model", old.Model, "--thinking", old.Thinking)
if restoreErr != nil || result.ExitCode != 0 { retErr = fmt.Errorf("%w; previous Pi settings could not be restored: recovery required", retErr) }
if json.Unmarshal([]byte(oldResult.Stdout), &old) != nil || old.Provider == "" || old.Model == "" || old.Thinking == "" {
return errors.New("Pi installation settings capture is invalid")
}
restore := func(cause error) error {
result, restoreErr := writeDefaults(context.Background(), runner, old)
if restoreErr != nil {
return fmt.Errorf("%w; previous Pi settings could not be restored: recovery required", cause)
}
}()
result, err := runCompose(ctx, runner, "exec", "-T", "core", "node", "/app/backend/dist/settings/settings-cli.js", "--provider", value.Provider, "--model", value.Model, "--thinking", value.Thinking)
if err != nil { return commandError("Pi installation settings write", result, err) }
wrote = true
if result.ExitCode != 0 {
return fmt.Errorf("%w; previous Pi settings could not be restored: recovery required", cause)
}
verified, readErr := readDefaults(context.Background(), runner, settingsArgs)
if readErr != nil || verified != old {
return fmt.Errorf("%w; previous Pi settings restoration could not be verified: recovery required", cause)
}
return cause
}
result, err := writeDefaults(ctx, runner, value)
if err != nil {
return restore(commandError("Pi installation settings write", result, err))
}
settings, err := runCompose(ctx, runner, settingsArgs...)
if err != nil { return commandError("Pi installation settings read-back", settings, err) }
if err != nil {
return restore(commandError("Pi installation settings read-back", settings, err))
}
var saved Defaults
if json.Unmarshal([]byte(settings.Stdout), &saved) != nil || saved.Provider != value.Provider || saved.Model != value.Model || saved.Thinking != value.Thinking {
return errors.New("Pi installation settings read-back did not match requested provider, model, and thinking")
return restore(errors.New("Pi installation settings read-back did not match requested provider, model, and thinking"))
}
after, err := renderedCore(ctx, runner)
if err != nil { return err }
if before.ConfigurationSHA != after.ConfigurationSHA { return errors.New("external endpoint configuration changed while configuring Pi") }
if err != nil {
return restore(err)
}
if before.ConfigurationSHA != after.ConfigurationSHA {
return restore(errors.New("external endpoint configuration changed while configuring Pi"))
}
return nil
}
func ConfigurationOptions(ctx context.Context, runner Runner) ([]ModelOption, error) {
args := append([]string{"exec", "-T", "core", "curl", "-fsS"}, internalIdentityHeaders...)
args = append(args, "http://127.0.0.1:8787/models")
models, err := runCompose(ctx, runner, args...)
if err != nil {
return nil, commandError("Pi options check", models, err)
}
var payload struct {
Models []ModelOption `json:"models"`
}
if json.Unmarshal([]byte(models.Stdout), &payload) != nil || len(payload.Models) == 0 {
return nil, errors.New("Pi options response is invalid or empty")
}
for _, option := range payload.Models {
if !choicePattern.MatchString(option.Provider) || !choicePattern.MatchString(option.ID) {
return nil, errors.New("Pi options response contains an invalid provider/model")
}
}
return payload.Models, nil
}
func writeDefaults(ctx context.Context, runner Runner, value Defaults) (compose.Result, error) {
return runCompose(ctx, runner, "exec", "-T", "core", "node", "/app/backend/dist/settings/settings-cli.js", "--provider", value.Provider, "--model", value.Model, "--thinking", value.Thinking)
}
func readDefaults(ctx context.Context, runner Runner, args []string) (Defaults, error) {
result, err := runCompose(ctx, runner, args...)
if err != nil {
return Defaults{}, commandError("Pi installation settings restoration read-back", result, err)
}
var value Defaults
if json.Unmarshal([]byte(result.Stdout), &value) != nil {
return Defaults{}, errors.New("Pi installation settings restoration read-back is invalid")
}
return value, nil
}
// Runner is the narrow, shell-free command boundary shared with thothctl.
type Runner interface {
Run(context.Context, []string, io.Reader) (compose.Result, error)
@@ -135,64 +186,70 @@ func Test(ctx context.Context, runner Runner) error {
if _, err := Status(ctx, runner); err != nil {
return err
}
for _, path := range []string{"health", "models", "settings"} {
args := []string{"exec", "-T", "core", "curl", "-fsS"}
if path != "health" { args = append(args, internalIdentityHeaders...) }
args = append(args, "http://127.0.0.1:8787/"+path)
result, err := runCompose(ctx, runner, args...)
if err != nil {
return commandError("Pi smoke check", result, err)
}
var payload any
if err := json.Unmarshal([]byte(result.Stdout), &payload); err != nil {
return fmt.Errorf("Pi smoke check returned invalid %s response", path)
}
object, ok := payload.(map[string]any)
if !ok {
return fmt.Errorf("Pi smoke check returned invalid %s response", path)
}
switch path {
case "health":
if object["status"] != "ok" { return errors.New("Pi smoke health response is not ready") }
case "models":
models, ok := object["models"].([]any)
if !ok || len(models) == 0 { return errors.New("Pi smoke models response is empty") }
valid := false
for _, item := range models { if model, ok := item.(map[string]any); ok && stringField(model, "provider") != "" && stringField(model, "id") != "" { valid = true; break } }
if !valid { return errors.New("Pi smoke models response has no provider/model choices") }
case "settings":
if stringField(object, "provider") == "" || stringField(object, "model") == "" || stringField(object, "thinking") == "" { return errors.New("Pi smoke settings response is incomplete") }
health, err := runCompose(ctx, runner, "exec", "-T", "core", "curl", "-fsS", "http://127.0.0.1:8787/health")
if err != nil {
return commandError("Pi smoke check", health, err)
}
var healthPayload struct {
Status string `json:"status"`
}
if json.Unmarshal([]byte(health.Stdout), &healthPayload) != nil || healthPayload.Status != "ok" {
return errors.New("Pi smoke health response is not ready")
}
models, err := ConfigurationOptions(ctx, runner)
if err != nil {
return err
}
settingsArgs := append([]string{"exec", "-T", "core", "curl", "-fsS"}, internalIdentityHeaders...)
settingsArgs = append(settingsArgs, "http://127.0.0.1:8787/settings")
settings, err := runCompose(ctx, runner, settingsArgs...)
if err != nil {
return commandError("Pi smoke settings check", settings, err)
}
var selected Defaults
if json.Unmarshal([]byte(settings.Stdout), &selected) != nil || !choicePattern.MatchString(selected.Provider) || !choicePattern.MatchString(selected.Model) || (selected.Thinking != "low" && selected.Thinking != "medium" && selected.Thinking != "high") {
return errors.New("Pi smoke settings response is incomplete")
}
for _, model := range models {
if model.Provider == selected.Provider && model.ID == selected.Model {
return nil
}
}
return nil
return errors.New("configured provider/model does not match an available Pi model entry")
}
func stringField(value map[string]any, key string) string { text, _ := value[key].(string); return strings.TrimSpace(text) }
func renderedCore(ctx context.Context, runner Runner) (Image, error) {
result, err := runCompose(ctx, runner, "config", "--format", "json")
if err != nil {
return Image{}, commandError("Compose configuration check", result, err)
}
var document struct {
Services map[string]struct {
Image string `json:"image"`
Environment map[string]any `json:"environment"`
} `json:"services"`
}
var document map[string]any
if err := json.Unmarshal([]byte(result.Stdout), &document); err != nil {
return Image{}, errors.New("Compose returned invalid rendered configuration")
}
core, exists := document.Services["core"]
if !exists || core.Image == "" {
services, ok := document["services"].(map[string]any)
if !ok {
return Image{}, errors.New("rendered Compose configuration has no services")
}
core, ok := services["core"].(map[string]any)
reference, _ := core["image"].(string)
if !ok || reference == "" {
return Image{}, errors.New("rendered Compose configuration has no core image")
}
endpoint, exists := core.Environment["THT_LLM_URL"].(string)
environment, _ := core["environment"].(map[string]any)
endpoint, exists := environment["THT_LLM_URL"].(string)
if !exists || strings.TrimSpace(endpoint) == "" {
return Image{}, errors.New("THT_LLM_URL must be configured before Pi lifecycle operations")
}
digest := sha256.Sum256([]byte(result.Stdout))
return Image{Reference: core.Image, ConfigurationSHA: fmt.Sprintf("%x", digest[:])}, nil
// Lifecycle overrides intentionally replace only core.image. Normalize that field so the
// non-secret configuration digest continues to detect endpoint/mount/configuration drift.
core["image"] = "<lifecycle-image>"
normalized, err := json.Marshal(document)
if err != nil {
return Image{}, errors.New("Compose configuration could not be normalized")
}
digest := sha256.Sum256(normalized)
return Image{Reference: reference, ConfigurationSHA: fmt.Sprintf("%x", digest[:])}, nil
}
func runCompose(ctx context.Context, runner Runner, args ...string) (compose.Result, error) {
@@ -2,8 +2,13 @@ package pi
import (
"context"
"encoding/json"
"errors"
"io"
"strings"
"testing"
"github.com/aritmolab/thothii/tools/thothctl/internal/compose"
)
func TestDoctorRequiresExternalEndpointAuthPiStateAndHealth(t *testing.T) {
@@ -16,6 +21,69 @@ func TestDoctorRequiresExternalEndpointAuthPiStateAndHealth(t *testing.T) {
}
}
func TestConfigureRestoresAndVerifiesOldSettingsAfterEveryPostSnapshotFailure(t *testing.T) {
for _, failure := range []string{"helper", "readback", "digest"} {
t.Run(failure, func(t *testing.T) {
fake := &configureRunner{failure: failure, settings: Defaults{Provider: "old", Model: "old-model", Thinking: "low"}}
err := Configure(context.Background(), fake, Defaults{Provider: "new", Model: "new-model", Thinking: "high"})
if err == nil {
t.Fatal("Configure() error = nil, want injected failure")
}
if fake.settings != (Defaults{Provider: "old", Model: "old-model", Thinking: "low"}) {
t.Fatalf("settings after failure = %#v, want old snapshot", fake.settings)
}
minimumReads := 3
if failure == "helper" {
minimumReads = 2
}
if fake.settingsReads < minimumReads {
t.Fatalf("settings read count = %d, want capture/failure reads plus verified restore", fake.settingsReads)
}
})
}
}
type configureRunner struct {
failure string
settings Defaults
settingsReads int
configReads int
}
func (f *configureRunner) Run(_ context.Context, args []string, _ io.Reader) (compose.Result, error) {
call := strings.Join(args, " ")
switch {
case strings.Contains(call, "config --format json"):
f.configReads++
endpoint := "https://llm.example.invalid"
if f.failure == "digest" && f.configReads > 1 {
endpoint = "https://drift.example.invalid"
}
return compose.Result{Stdout: `{"services":{"core":{"image":"thothii-core:local","environment":{"THT_LLM_URL":"` + endpoint + `"}}}}`}, nil
case strings.Contains(call, "/models"):
return compose.Result{Stdout: `{"models":[{"provider":"old","id":"old-model"},{"provider":"new","id":"new-model"}]}`}, nil
case strings.Contains(call, "settings-cli.js"):
if strings.Contains(call, "--provider new") {
f.settings = Defaults{Provider: "new", Model: "new-model", Thinking: "high"}
if f.failure == "helper" {
return compose.Result{ExitCode: 17}, errors.New("injected helper failure")
}
} else {
f.settings = Defaults{Provider: "old", Model: "old-model", Thinking: "low"}
}
return compose.Result{}, nil
case strings.Contains(call, "/settings"):
f.settingsReads++
if f.failure == "readback" && f.settingsReads == 2 {
return compose.Result{Stdout: `{}`}, nil
}
contents, _ := json.Marshal(f.settings)
return compose.Result{Stdout: string(contents)}, nil
default:
return compose.Result{}, nil
}
}
func TestConfigureValidatesBackendModelOptionsWritesRealCoreSettingsAndUsesUpstreamIdentity(t *testing.T) {
fake := newFakeRunner()
if err := Configure(context.Background(), fake, Defaults{Provider: "provider", Model: "model", Thinking: "medium"}); err != nil {
@@ -43,3 +111,16 @@ func TestTestUsesOnlySanitizedPiAndCoreProbes(t *testing.T) {
t.Fatalf("probe commands expose secret: %s", got)
}
}
func TestTestRequiresConfiguredProviderAndModelToMatchOneAvailableEntry(t *testing.T) {
fake := newFakeRunner()
fake.modelsWire = `{"models":[{"id":"different-model","provider":"provider"}]}`
if err := Test(context.Background(), fake); err == nil || !strings.Contains(err.Error(), "configured provider/model") {
t.Fatalf("Test() error = %v, want exact settings/model mismatch", err)
}
fake.modelsWire = `{"models":[{"id":"model","provider":"provider"}]}`
if err := Test(context.Background(), fake); err != nil {
t.Fatalf("Test() exact match error = %v", err)
}
assertCalled(t, fake.calls, "pi --version")
}
+23 -3
View File
@@ -2,14 +2,34 @@
package pi
import "os"
import (
"errors"
"os"
"path/filepath"
)
// durableReplace acknowledges both the data file and its directory entry. A successful return
// is the strongest atomic replacement guarantee supported by Unix filesystems.
func durableReplace(temporary, target, directory string) error {
if err := os.Rename(temporary, target); err != nil { return err }
if err := os.Rename(temporary, target); err != nil {
return err
}
dir, err := os.Open(directory)
if err != nil { return err }
if err != nil {
return err
}
defer dir.Close()
return dir.Sync()
}
func durableRemove(path string) error {
if err := os.Remove(path); err != nil && !errors.Is(err, os.ErrNotExist) {
return err
}
dir, err := os.Open(filepath.Dir(path))
if err != nil {
return err
}
defer dir.Close()
return dir.Sync()
}
+19 -3
View File
@@ -2,14 +2,30 @@
package pi
import "golang.org/x/sys/windows"
import (
"errors"
"os"
"golang.org/x/sys/windows"
)
// MoveFileEx requests replacement and write-through on Windows. Directory fsync is not exposed
// by the Windows API in the same form as Unix, so callers must not claim a stronger guarantee.
func durableReplace(temporary, target, _ string) error {
from, err := windows.UTF16PtrFromString(temporary)
if err != nil { return err }
if err != nil {
return err
}
to, err := windows.UTF16PtrFromString(target)
if err != nil { return err }
if err != nil {
return err
}
return windows.MoveFileEx(from, to, windows.MOVEFILE_REPLACE_EXISTING|windows.MOVEFILE_WRITE_THROUGH)
}
func durableRemove(path string) error {
if err := os.Remove(path); err != nil && !errors.Is(err, os.ErrNotExist) {
return err
}
return nil
}
@@ -1,16 +0,0 @@
//go:build !windows
package pi
import (
"errors"
"os"
"syscall"
)
func processAlive(pid int) bool {
process, err := os.FindProcess(pid)
if err != nil { return false }
err = process.Signal(syscall.Signal(0))
return err == nil || errors.Is(err, syscall.EPERM)
}
@@ -1,14 +0,0 @@
//go:build windows
package pi
import "golang.org/x/sys/windows"
func processAlive(pid int) bool {
handle, err := windows.OpenProcess(windows.PROCESS_QUERY_LIMITED_INFORMATION, false, uint32(pid))
if err != nil { return err == windows.ERROR_ACCESS_DENIED }
defer windows.CloseHandle(handle)
var code uint32
if windows.GetExitCodeProcess(handle, &code) != nil { return true }
return code == 259 // STILL_ACTIVE
}
+67 -61
View File
@@ -2,18 +2,20 @@
package pi
import (
"crypto/sha256"
"encoding/json"
"errors"
"fmt"
"crypto/sha256"
"os"
"path/filepath"
"sort"
"strings"
"time"
"github.com/gofrs/flock"
)
const stateFileVersion = 2
const stateFileVersion = 3
// Phase describes the durable point reached by a Pi update.
type Phase string
@@ -30,22 +32,21 @@ const (
// Image is the non-secret recovery identity of a core image and its mounted volume names.
type Image struct {
ID string `json:"id"`
Reference string `json:"reference"`
Volumes []string `json:"volumes"`
Mounts []Mount `json:"mounts"`
MountFingerprint string `json:"mount_fingerprint"`
ConfigurationSHA string `json:"configuration_sha256,omitempty"`
ID string `json:"id"`
Reference string `json:"reference"`
Mounts []Mount `json:"mounts"`
MountFingerprint string `json:"mount_fingerprint"`
ConfigurationSHA string `json:"configuration_sha256,omitempty"`
}
// Mount is the complete persistence identity relevant to safe core recreation.
type Mount struct {
Type string `json:"type"`
Name string `json:"name,omitempty"`
Type string `json:"type"`
Name string `json:"name,omitempty"`
SourceSHA256 string `json:"source_sha256"`
Destination string `json:"destination"`
RW bool `json:"rw"`
Options string `json:"options,omitempty"`
Destination string `json:"destination"`
RW bool `json:"rw"`
Options string `json:"options,omitempty"`
}
// Target records the immutable input selected by the operator. Source is either build or a
@@ -58,13 +59,14 @@ type Target struct {
// State is recovery metadata stored below the installation project. It never stores environment
// values, secret paths, credentials, or command output.
type State struct {
Version int `json:"version"`
Phase Phase `json:"phase"`
UpdatedAt time.Time `json:"updated_at"`
Target Target `json:"target,omitempty"`
Previous Image `json:"previous"`
Candidate Image `json:"candidate,omitempty"`
Error string `json:"error,omitempty"`
Version int `json:"version"`
Transaction string `json:"transaction"`
Phase Phase `json:"phase"`
UpdatedAt time.Time `json:"updated_at"`
Target Target `json:"target,omitempty"`
Previous Image `json:"previous"`
Candidate Image `json:"candidate,omitempty"`
Error string `json:"error,omitempty"`
}
func readState(path string) (State, error) {
@@ -104,15 +106,30 @@ func writeState(path string, state State) error {
func writeFileDurably(path, prefix string, contents []byte) error {
directory := filepath.Dir(path)
if err := os.MkdirAll(directory, 0o700); err != nil { return err }
if err := os.MkdirAll(directory, 0o700); err != nil {
return err
}
temporary, err := os.CreateTemp(directory, prefix+"*.tmp")
if err != nil { return err }
if err != nil {
return err
}
temporaryName := temporary.Name()
defer os.Remove(temporaryName)
if err := temporary.Chmod(0o600); err != nil { temporary.Close(); return err }
if _, err := temporary.Write(contents); err != nil { temporary.Close(); return err }
if err := temporary.Sync(); err != nil { temporary.Close(); return err }
if err := temporary.Close(); err != nil { return err }
if err := temporary.Chmod(0o600); err != nil {
temporary.Close()
return err
}
if _, err := temporary.Write(contents); err != nil {
temporary.Close()
return err
}
if err := temporary.Sync(); err != nil {
temporary.Close()
return err
}
if err := temporary.Close(); err != nil {
return err
}
return durableReplace(temporaryName, path, directory)
}
@@ -138,7 +155,10 @@ type lockOwner struct {
Transaction string `json:"transaction"`
}
type updateLock struct{ path string }
type updateLock struct {
file *flock.Flock
metadata string
}
var ErrLockHeld = errors.New("another Pi update or rollback is already in progress")
@@ -147,47 +167,33 @@ func acquireLock(statePath string) (*updateLock, error) {
return nil, errors.New("could not create Pi update recovery directory")
}
path := statePath + ".lock"
file, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o600)
file := flock.New(path, flock.SetPermissions(0o600))
locked, err := file.TryLock()
if err != nil {
if errors.Is(err, os.ErrExist) {
if reclaimDeadLocalLock(path, statePath) {
return acquireLock(statePath)
}
return nil, ErrLockHeld
}
return nil, errors.New("could not acquire Pi update lock")
}
if !locked {
return nil, ErrLockHeld
}
host, err := os.Hostname()
if err != nil { _ = file.Close(); _ = os.Remove(path); return nil, errors.New("could not identify Pi update lock owner") }
if err != nil {
_ = file.Unlock()
return nil, errors.New("could not identify Pi update lock owner")
}
owner := lockOwner{PID: os.Getpid(), Host: host, StartedAt: time.Now().UTC(), Transaction: fmt.Sprintf("%d-%d", os.Getpid(), time.Now().UnixNano())}
contents, err := json.Marshal(owner)
if err != nil { _ = file.Close(); _ = os.Remove(path); return nil, errors.New("could not record Pi update lock owner") }
if _, err := file.Write(append(contents, '\n')); err != nil || file.Sync() != nil || file.Close() != nil {
_ = file.Close(); _ = os.Remove(path)
if err != nil {
_ = file.Unlock()
return nil, errors.New("could not record Pi update lock owner")
}
return &updateLock{path: path}, nil
metadata := path + ".owner.json"
if err := writeFileDurably(metadata, ".lock-owner-", append(contents, '\n')); err != nil {
_ = file.Unlock()
return nil, errors.New("could not record Pi update lock owner")
}
return &updateLock{file: file, metadata: metadata}, nil
}
func (l *updateLock) Release() { _ = os.Remove(l.path) }
// reclaimDeadLocalLock is deliberately conservative: a malformed, remote, or merely old lock
// is recovery-required. Only a process we can prove is gone on this machine is reclaimed.
func reclaimDeadLocalLock(path, statePath string) bool {
info, err := os.Stat(path)
if err != nil || time.Since(info.ModTime()) < 5*time.Minute || !hasPendingRecoveryState(statePath) { return false }
contents, err := os.ReadFile(path)
if err != nil { return os.Remove(path) == nil }
var owner lockOwner
if json.Unmarshal(contents, &owner) != nil || owner.PID <= 0 || owner.Host == "" { return false }
host, err := os.Hostname()
if err != nil || owner.Host != host { return false }
if processAlive(owner.PID) { return false }
return os.Remove(path) == nil
}
func hasPendingRecoveryState(path string) bool {
contents, err := os.ReadFile(path); if err != nil { return false }
var state State
if json.Unmarshal(contents, &state) != nil { return false }
return state.Phase != PhaseVerified && state.Phase != PhaseRolledBack && state.Phase != PhaseNoop
func (l *updateLock) Release() {
_ = durableRemove(l.metadata)
_ = l.file.Unlock()
}
+61
View File
@@ -0,0 +1,61 @@
package pi
import (
"errors"
"os"
"os/exec"
"path/filepath"
"testing"
)
func TestAdvisoryLockRejectsAConcurrentOwner(t *testing.T) {
statePath := filepath.Join(t.TempDir(), "update-state.json")
first, err := acquireLock(statePath)
if err != nil {
t.Fatal(err)
}
defer first.Release()
second, err := acquireLock(statePath)
if second != nil {
second.Release()
}
if !errors.Is(err, ErrLockHeld) {
t.Fatalf("second acquireLock() error = %v, want ErrLockHeld", err)
}
}
func TestAdvisoryLockCrashReleasesAndReacquires(t *testing.T) {
statePath := filepath.Join(t.TempDir(), "update-state.json")
if os.Getenv("THOTHCTL_LOCK_CRASH_HELPER") == "1" {
lock, err := acquireLock(os.Getenv("THOTHCTL_LOCK_STATE"))
if err != nil || lock == nil {
os.Exit(23)
}
os.Exit(0) // Deliberately bypass Release: the OS must release ownership.
}
command := exec.Command(os.Args[0], "-test.run=^TestAdvisoryLockCrashReleasesAndReacquires$")
command.Env = append(os.Environ(), "THOTHCTL_LOCK_CRASH_HELPER=1", "THOTHCTL_LOCK_STATE="+statePath)
if output, err := command.CombinedOutput(); err != nil {
t.Fatalf("crash helper failed: %v: %s", err, output)
}
lock, err := acquireLock(statePath)
if err != nil {
t.Fatalf("acquireLock() after owner crash = %v", err)
}
lock.Release()
}
func TestAdvisoryLockIgnoresPartialDiagnosticMetadata(t *testing.T) {
statePath := filepath.Join(t.TempDir(), "update-state.json")
if err := os.WriteFile(statePath+".lock", nil, 0o600); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(statePath+".lock.owner.json", []byte("{partial"), 0o600); err != nil {
t.Fatal(err)
}
lock, err := acquireLock(statePath)
if err != nil {
t.Fatalf("acquireLock() with partial diagnostics = %v", err)
}
lock.Release()
}
+329 -98
View File
@@ -2,15 +2,19 @@ package pi
import (
"context"
"crypto/sha256"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"regexp"
"sort"
"strings"
"time"
"github.com/aritmolab/thothii/tools/thothctl/internal/compose"
"github.com/distribution/reference"
)
@@ -46,29 +50,49 @@ type Result struct {
StatePath string
}
type lifecycleHooks struct {
writeState func(string, State) error
removeFile func(string) error
sleep func(time.Duration)
}
var defaultLifecycleHooks = lifecycleHooks{
writeState: writeState,
removeFile: durableRemove,
sleep: time.Sleep,
}
// Update performs a recoverable core-only Pi update using the default Compose command layout.
func Update(ctx context.Context, runner Runner, request Request) (result Result, retErr error) {
return updateWithHooks(ctx, runner, request, defaultLifecycleHooks)
}
func updateWithHooks(ctx context.Context, runner Runner, request Request, hooks lifecycleHooks) (result Result, retErr error) {
if request.StatePath == "" {
return Result{}, errors.New("update state path is required")
}
lock, err := acquireLock(request.StatePath)
if err != nil {
return Result{StatePath: request.StatePath}, err
}
defer lock.Release()
if request.StatePath == "" {
return Result{}, errors.New("update state path is required")
}
if !request.Confirm {
return Result{StatePath: request.StatePath}, ErrConfirmationRequired
}
if !versionPattern.MatchString(request.Version) {
return Result{StatePath: request.StatePath}, fmt.Errorf("%w: Pi version must be an explicit pinned version", ErrInvalidRequest)
}
if request.Source == "" { return Result{StatePath: request.StatePath}, fmt.Errorf("%w: Pi update requires an explicit source: build or pull", ErrInvalidRequest) }
if request.Source == "" {
return Result{StatePath: request.StatePath}, fmt.Errorf("%w: Pi update requires an explicit source: build or pull", ErrInvalidRequest)
}
if request.Source != BuildSource && request.Source != PullSource {
return Result{StatePath: request.StatePath}, fmt.Errorf("%w: Pi update source must be build or pull", ErrInvalidRequest)
}
if request.Source == PullSource {
canonical, err := canonicalDigestReference(request.Image)
if err != nil { return Result{StatePath: request.StatePath}, fmt.Errorf("%w: %v", ErrInvalidRequest, err) }
if err != nil {
return Result{StatePath: request.StatePath}, fmt.Errorf("%w: %v", ErrInvalidRequest, err)
}
request.Image = canonical
}
if old, err := readState(request.StatePath); err == nil && old.Phase != PhaseVerified && old.Phase != PhaseRolledBack && old.Phase != PhaseNoop {
@@ -76,12 +100,21 @@ func Update(ctx context.Context, runner Runner, request Request) (result Result,
} else if err != nil && !errors.Is(err, os.ErrNotExist) {
return Result{StatePath: request.StatePath}, err
}
if err := setMaintenance(ctx, runner, true); err != nil { return Result{StatePath: request.StatePath}, err }
if err := setMaintenance(ctx, runner, true); err != nil {
return Result{StatePath: request.StatePath}, err
}
clearMaintenance := true
defer func() {
if !clearMaintenance {
return
}
if clearErr := setMaintenance(context.Background(), runner, false); clearErr != nil {
result = Result{Phase: PhaseFailed, StatePath: request.StatePath}
if retErr == nil { retErr = errors.New("maintenance admission gate could not be cleared: recovery required")
} else { retErr = fmt.Errorf("%w; maintenance admission gate could not be cleared: recovery required", retErr) }
if retErr == nil {
retErr = errors.New("maintenance admission gate could not be cleared: recovery required")
} else {
retErr = fmt.Errorf("%w; maintenance admission gate could not be cleared: recovery required", retErr)
}
}
}()
@@ -101,7 +134,7 @@ func Update(ctx context.Context, runner Runner, request Request) (result Result,
if !running {
break
}
time.Sleep(time.Second)
hooks.sleep(time.Second)
}
if running {
return Result{StatePath: request.StatePath}, ErrActiveSessions
@@ -127,46 +160,87 @@ func Update(ctx context.Context, runner Runner, request Request) (result Result,
return Result{StatePath: request.StatePath}, err
}
previous.ConfigurationSHA = configured.ConfigurationSHA
state := State{Phase: PhasePreflight, Target: Target{Version: request.Version, Source: sourceValue(request)}, Previous: previous}
if err := writeState(request.StatePath, state); err != nil {
transaction := lifecycleTransaction(request.StatePath)
previous.Reference = lifecycleImageTag(transaction, "previous")
candidateReference := lifecycleImageTag(transaction, "candidate")
if err := tagImage(ctx, runner, previous.ID, previous.Reference, "previous Pi image pin"); err != nil {
return Result{StatePath: request.StatePath}, err
}
state := State{
Transaction: transaction,
Phase: PhasePreflight,
Target: Target{Version: request.Version, Source: sourceValue(request)},
Previous: previous,
Candidate: Image{Reference: candidateReference},
}
if err := hooks.writeState(request.StatePath, state); err != nil {
return Result{StatePath: request.StatePath}, err
}
overridePath := lifecycleOverridePath(request.StatePath, transaction)
if err := writeLifecycleOverride(overridePath, candidateReference); err != nil {
return Result{StatePath: request.StatePath}, err
}
lifecycle := composeOverrideRunner{Runner: runner, path: overridePath}
state.Phase = PhaseBuilding
if err := writeState(request.StatePath, state); err != nil { return Result{Phase: PhaseFailed, StatePath: request.StatePath}, err }
if err := prepareCandidate(ctx, runner, request, previous.Reference); err != nil {
return compensate(ctx, runner, request.StatePath, state, err)
clearMaintenance = false
if err := hooks.writeState(request.StatePath, state); err != nil {
result, retErr, clearMaintenance = compensate(ctx, runner, request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
if request.Drain {
running, err = activeSessions(ctx, runner)
if err != nil {
return compensate(ctx, runner, request.StatePath, state, err)
}
if running {
return compensate(ctx, runner, request.StatePath, state, ErrActiveSessions)
}
if err := prepareCandidate(ctx, lifecycle, request, candidateReference); err != nil {
result, retErr, clearMaintenance = compensate(ctx, runner, request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
if err := recreateCore(ctx, runner); err != nil {
return compensate(ctx, runner, request.StatePath, state, err)
running, err = activeSessions(ctx, runner)
if err != nil {
result, retErr, clearMaintenance = compensate(ctx, runner, request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
if running {
result, retErr, clearMaintenance = compensate(ctx, runner, request.StatePath, overridePath, state, ErrActiveSessions, hooks)
return result, retErr
}
if err := recreateCore(ctx, lifecycle); err != nil {
result, retErr, clearMaintenance = compensate(ctx, runner, request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
if err := ensureMaintenance(ctx, lifecycle); err != nil {
result, retErr, clearMaintenance = compensate(ctx, runner, request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
state.Phase = PhaseRecreated
state.Candidate, err = runningImage(ctx, runner, previous.Reference)
if err != nil { return compensate(ctx, runner, request.StatePath, state, err) }
if err := writeState(request.StatePath, state); err != nil {
return compensate(ctx, runner, request.StatePath, state, err)
state.Candidate, err = runningImage(ctx, lifecycle, candidateReference)
if err != nil {
result, retErr, clearMaintenance = compensate(ctx, runner, request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
if err := verifyCandidate(ctx, runner, request.Version, previous); err != nil {
return compensate(ctx, runner, request.StatePath, state, err)
if err := hooks.writeState(request.StatePath, state); err != nil {
result, retErr, clearMaintenance = compensate(ctx, runner, request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
if err := verifyCandidate(ctx, lifecycle, request.Version, previous); err != nil {
result, retErr, clearMaintenance = compensate(ctx, runner, request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
state.Phase, state.Error = PhaseVerified, ""
if err := writeState(request.StatePath, state); err != nil {
return compensate(ctx, runner, request.StatePath, state, err)
if err := hooks.writeState(request.StatePath, state); err != nil {
result, retErr, clearMaintenance = compensate(ctx, runner, request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
if err := hooks.removeFile(overridePath); err != nil {
return Result{Phase: PhaseFailed, StatePath: request.StatePath}, errors.New("verified update override cleanup failed: maintenance recovery required")
}
clearMaintenance = true
return Result{Phase: PhaseVerified, StatePath: request.StatePath}, nil
}
// 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 rollbackWithHooks(ctx context.Context, runner Runner, statePath string, confirm bool, hooks lifecycleHooks) (result Result, retErr error) {
lock, err := acquireLock(statePath)
if err != nil {
return Result{StatePath: statePath}, err
@@ -175,51 +249,88 @@ func Rollback(ctx context.Context, runner Runner, statePath string, confirm bool
if !confirm {
return Result{StatePath: statePath}, ErrConfirmationRequired
}
if err := setMaintenance(ctx, runner, true); err != nil { return Result{StatePath: statePath}, err }
defer func() {
if clearErr := setMaintenance(context.Background(), runner, false); clearErr != nil {
result = Result{Phase: PhaseFailed, StatePath: statePath}
if retErr == nil { retErr = errors.New("maintenance admission gate could not be cleared: recovery required")
} else { retErr = fmt.Errorf("%w; maintenance admission gate could not be cleared: recovery required", retErr) }
}
}()
if active, err := activeSessions(ctx, runner); err != nil { return Result{StatePath: statePath}, err
} else if active { return Result{StatePath: statePath}, ErrActiveSessions }
state, err := readState(statePath)
if err != nil {
if err := setMaintenance(ctx, runner, true); err != nil {
return Result{StatePath: statePath}, err
}
if err := restore(ctx, runner, state.Previous); err != nil {
clearMaintenance := true
defer func() {
if !clearMaintenance {
return
}
if clearErr := setMaintenance(context.Background(), runner, false); clearErr != nil {
result = Result{Phase: PhaseFailed, StatePath: statePath}
if retErr == nil {
retErr = errors.New("maintenance admission gate could not be cleared: recovery required")
} else {
retErr = fmt.Errorf("%w; maintenance admission gate could not be cleared: recovery required", retErr)
}
}
}()
if active, err := activeSessions(ctx, runner); err != nil {
return Result{StatePath: statePath}, err
} else if active {
return Result{StatePath: statePath}, ErrActiveSessions
}
state, err := readState(statePath)
if err != nil {
clearMaintenance = false
return Result{StatePath: statePath}, err
}
overridePath := lifecycleOverridePath(statePath, state.Transaction)
if err := writeLifecycleOverride(overridePath, state.Previous.Reference); err != nil {
clearMaintenance = false
return Result{StatePath: statePath}, err
}
clearMaintenance = false
lifecycle := composeOverrideRunner{Runner: runner, path: overridePath}
if err := restore(ctx, lifecycle, state.Previous); err != nil {
state.Phase, state.Error = PhaseFailed, "rollback failed"
if writeErr := writeState(statePath, state); writeErr != nil { return Result{Phase: PhaseFailed, StatePath: statePath}, errors.New("rollback failed and recovery state could not be persisted") }
if writeErr := hooks.writeState(statePath, state); writeErr != nil {
return Result{Phase: PhaseFailed, StatePath: statePath}, errors.New("rollback failed and recovery state could not be persisted")
}
return Result{Phase: PhaseFailed, StatePath: statePath}, err
}
state.Phase, state.Error = PhaseRolledBack, ""
if err := writeState(statePath, state); err != nil {
if err := hooks.writeState(statePath, state); err != nil {
return Result{Phase: PhaseFailed, StatePath: statePath}, errors.New("rollback restored the core but recovery state could not be persisted")
}
if err := hooks.removeFile(overridePath); err != nil {
return Result{Phase: PhaseFailed, StatePath: statePath}, errors.New("rollback override cleanup failed: maintenance recovery required")
}
clearMaintenance = true
return Result{Phase: PhaseRolledBack, StatePath: statePath}, nil
}
func compensate(ctx context.Context, runner Runner, statePath string, state State, cause error) (Result, error) {
if restoreErr := restore(ctx, runner, state.Previous); restoreErr != nil {
func compensate(ctx context.Context, runner Runner, statePath, overridePath string, state State, cause error, hooks lifecycleHooks) (Result, error, bool) {
if err := ensureMaintenance(context.Background(), runner); err != nil {
state.Phase, state.Error = PhaseFailed, "maintenance recovery failed"
_ = hooks.writeState(statePath, state)
return Result{Phase: PhaseFailed, StatePath: statePath}, errors.New("update failed and maintenance could not be reactivated: recovery required"), false
}
if active, err := activeSessions(context.Background(), runner); err != nil || active {
state.Phase, state.Error = PhaseFailed, "rollback inventory failed"
_ = hooks.writeState(statePath, state)
return Result{Phase: PhaseFailed, StatePath: statePath}, errors.New("update failed and rollback inventory is not quiescent: recovery required"), false
}
if err := writeLifecycleOverride(overridePath, state.Previous.Reference); err != nil {
return Result{Phase: PhaseFailed, StatePath: statePath}, errors.New("update failed and rollback override could not be prepared: recovery required"), false
}
lifecycle := composeOverrideRunner{Runner: runner, path: overridePath}
if restoreErr := restore(ctx, lifecycle, state.Previous); restoreErr != nil {
state.Phase, state.Error = PhaseFailed, "candidate verification and automatic rollback failed"
if writeErr := writeState(statePath, state); writeErr != nil {
return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("update failed and rollback proof failed; recovery state could not be persisted")
if writeErr := hooks.writeState(statePath, state); writeErr != nil {
return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("update failed and rollback proof failed; recovery state could not be persisted"), false
}
return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("update failed; automatic rollback also failed: recovery required")
return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("update failed; automatic rollback also failed: recovery required"), false
}
state.Phase, state.Error = PhaseRolledBack, ""
if writeErr := writeState(statePath, state); writeErr != nil {
return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("previous core image was restored but recovery state write failed: recovery required")
if writeErr := hooks.writeState(statePath, state); writeErr != nil {
return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("previous core image was restored but recovery state write failed: recovery required"), false
}
return Result{Phase: PhaseRolledBack, StatePath: statePath}, fmt.Errorf("update failed; previous core image was restored")
}
func recordFailure(path string, state State, label string, cause error) error {
state.Phase, state.Error = PhaseFailed, label
if err := writeState(path, state); err != nil { return fmt.Errorf("%w; recovery state write failed", cause) }
return cause
if err := hooks.removeFile(overridePath); err != nil {
return Result{Phase: PhaseFailed, StatePath: statePath}, errors.New("previous core image was restored but override cleanup failed: recovery required"), false
}
return Result{Phase: PhaseRolledBack, StatePath: statePath}, fmt.Errorf("update failed; previous core image was restored"), true
}
func sourceValue(request Request) string {
@@ -234,7 +345,9 @@ func canonicalDigestReference(value string) (string, error) {
return "", errors.New("pulled Pi image must be a credential-free canonical sha256 digest reference")
}
parsed, err := reference.ParseAnyReference(value)
if err != nil { return "", errors.New("pulled Pi image must be a valid canonical sha256 digest reference") }
if err != nil {
return "", errors.New("pulled Pi image must be a valid canonical sha256 digest reference")
}
canonical, ok := parsed.(reference.Canonical)
if !ok || canonical.Digest().Algorithm().String() != "sha256" || len(canonical.Digest().Encoded()) != 64 {
return "", errors.New("pulled Pi image must use an immutable sha256 digest")
@@ -244,14 +357,56 @@ func canonicalDigestReference(value string) (string, error) {
func setMaintenance(ctx context.Context, runner Runner, enabled bool) error {
path := "deactivate"
if enabled { path = "activate" }
args := append([]string{"exec", "-T", "core", "curl", "-fsS", "-X", "POST"}, internalIdentityHeaders...)
args = append(args, "http://127.0.0.1:8787/internal/maintenance/"+path)
if enabled {
path = "activate"
}
args := []string{"exec", "-T", "core", "curl", "-fsS", "-X", "POST", "http://127.0.0.1:8787/internal/maintenance/" + path}
result, err := runCompose(ctx, runner, args...)
if err != nil { return commandError("maintenance admission gate", result, err) }
var status struct { Active bool `json:"active"`; Admissions int `json:"admissions"` }
if json.Unmarshal([]byte(result.Stdout), &status) != nil || status.Active != enabled || status.Admissions != 0 { return errors.New("maintenance admission gate did not acknowledge a quiescent state") }
return nil
status, valid := parseMaintenanceStatus(result.Stdout)
if err == nil && valid && status.Active == enabled && status.Admissions == 0 {
return nil
}
// A core recreate or transport interruption may lose only the response. Resolve ambiguity by
// reading the durable gate state before deciding that operator recovery is required.
observed, statusErr := MaintenanceStatus(ctx, runner)
if statusErr == nil && observed.Active == enabled && observed.Admissions == 0 {
return nil
}
if err != nil {
return commandError("maintenance admission gate", result, err)
}
return errors.New("maintenance admission gate did not acknowledge a quiescent state")
}
type MaintenanceState struct {
Active bool `json:"active"`
Admissions int `json:"admissions"`
}
func parseMaintenanceStatus(value string) (MaintenanceState, bool) {
var status MaintenanceState
err := json.Unmarshal([]byte(value), &status)
return status, err == nil && status.Admissions >= 0
}
func MaintenanceStatus(ctx context.Context, runner Runner) (MaintenanceState, error) {
result, err := runCompose(ctx, runner, "exec", "-T", "core", "curl", "-fsS", "http://127.0.0.1:8787/internal/maintenance/status")
if err != nil {
return MaintenanceState{}, commandError("maintenance status check", result, err)
}
status, valid := parseMaintenanceStatus(result.Stdout)
if !valid {
return MaintenanceState{}, errors.New("maintenance status check returned invalid data")
}
return status, nil
}
func ensureMaintenance(ctx context.Context, runner Runner) error {
status, err := MaintenanceStatus(ctx, runner)
if err == nil && status.Active && status.Admissions == 0 {
return nil
}
return setMaintenance(ctx, runner, true)
}
func activeSessions(ctx context.Context, runner Runner) (bool, error) {
@@ -261,14 +416,14 @@ func activeSessions(ctx context.Context, runner Runner) (bool, error) {
if err != nil {
return false, commandError("active-session check", result, err)
}
var payload struct { Sessions []struct {
var payload []struct {
Status string `json:"status"`
Archived bool `json:"archived"`
} `json:"sessions"` }
}
if err := json.Unmarshal([]byte(result.Stdout), &payload); err != nil {
return false, errors.New("active-session check returned invalid session data")
}
for _, session := range payload.Sessions {
for _, session := range payload {
if !session.Archived && session.Status != "finalized" && session.Status != "closed" {
return true, nil
}
@@ -291,9 +446,14 @@ func runningImage(ctx context.Context, runner Runner, reference string) (Image,
return Image{}, commandError("core volume check", mounts, err)
}
var raw []struct {
Type string `json:"Type"`; Name string `json:"Name"`; Source string `json:"Source"`
Destination string `json:"Destination"`; RW bool `json:"RW"`; Mode string `json:"Mode"`
Propagation string `json:"Propagation"`; Driver string `json:"Driver"`
Type string `json:"Type"`
Name string `json:"Name"`
Source string `json:"Source"`
Destination string `json:"Destination"`
RW bool `json:"RW"`
Mode string `json:"Mode"`
Propagation string `json:"Propagation"`
Driver string `json:"Driver"`
}
if err := json.Unmarshal([]byte(mounts.Stdout), &raw); err != nil {
return Image{}, errors.New("core returned invalid persistence mount data")
@@ -302,18 +462,16 @@ func runningImage(ctx context.Context, runner Runner, reference string) (Image,
return Image{}, errors.New("core has no persistence mounts to preserve")
}
contract := make([]Mount, 0, len(raw))
volumes := make([]string, 0, len(raw))
for _, mount := range raw {
if mount.Type == "" || mount.Source == "" || mount.Destination == "" { return Image{}, errors.New("core returned incomplete persistence mount data") }
contract = append(contract, Mount{Type: mount.Type, Name: mount.Name, SourceSHA256: mountSourceHash(mount.Source), Destination: mount.Destination, RW: mount.RW, Options: strings.Join([]string{mount.Mode, mount.Propagation, mount.Driver}, "\x00")})
if mount.Type == "volume" && mount.Name != "" {
volumes = append(volumes, mount.Name)
if mount.Type == "" || mount.Source == "" || mount.Destination == "" {
return Image{}, errors.New("core returned incomplete persistence mount data")
}
contract = append(contract, Mount{Type: mount.Type, Name: mount.Name, SourceSHA256: mountSourceHash(mount.Source), Destination: mount.Destination, RW: mount.RW, Options: strings.Join([]string{mount.Mode, mount.Propagation, mount.Driver}, "\x00")})
}
return Image{ID: strings.TrimSpace(image.Stdout), Reference: reference, Volumes: volumes, Mounts: contract, MountFingerprint: mountFingerprint(contract)}, nil
return Image{ID: strings.TrimSpace(image.Stdout), Reference: reference, Mounts: contract, MountFingerprint: mountFingerprint(contract)}, nil
}
func prepareCandidate(ctx context.Context, runner Runner, request Request, reference string) error {
func prepareCandidate(ctx context.Context, runner Runner, request Request, candidateReference string) error {
if request.Source == BuildSource {
result, err := runCompose(ctx, runner, "build", "--pull", "--build-arg", "PI_VERSION="+request.Version, "core")
if err != nil {
@@ -325,11 +483,7 @@ func prepareCandidate(ctx context.Context, runner Runner, request Request, refer
if err != nil {
return commandError("Pi image pull", pull, err)
}
tag, err := runner.Run(ctx, []string{"image", "tag", request.Image, reference}, nil)
if err != nil {
return commandError("Pi image tag", tag, err)
}
return nil
return tagImage(ctx, runner, request.Image, candidateReference, "Pi image tag")
}
func recreateCore(ctx context.Context, runner Runner) error {
@@ -373,13 +527,15 @@ func verifyCandidate(ctx context.Context, runner Runner, wanted string, previous
}
func restore(ctx context.Context, runner Runner, previous Image) error {
tag, err := runner.Run(ctx, []string{"image", "tag", previous.ID, previous.Reference}, nil)
if err != nil {
return commandError("rollback image restore", tag, err)
if err := tagImage(ctx, runner, previous.ID, previous.Reference, "rollback image restore"); err != nil {
return err
}
if err := recreateCore(ctx, runner); err != nil {
return err
}
if err := ensureMaintenance(ctx, runner); err != nil {
return err
}
configured, err := renderedCore(ctx, runner)
if err != nil {
return err
@@ -406,16 +562,91 @@ func restore(ctx context.Context, runner Runner, previous Image) error {
return nil
}
func nonEmptyLines(text string) []string {
var values []string
for _, value := range strings.Split(text, "\n") {
if value = strings.TrimSpace(value); value != "" {
values = append(values, value)
}
func tagImage(ctx context.Context, runner Runner, source, target, label string) error {
result, err := runner.Run(ctx, []string{"image", "tag", source, target}, nil)
if err != nil {
return commandError(label, result, err)
}
sort.Strings(values)
return values
return nil
}
type composeOverrideRunner struct {
Runner
path string
}
func (r composeOverrideRunner) Run(ctx context.Context, args []string, stdin io.Reader) (compose.Result, error) {
if len(args) > 0 && args[0] == "compose" {
withOverride := append([]string{"compose", "-f", r.path}, args[1:]...)
return r.Runner.Run(ctx, withOverride, stdin)
}
return r.Runner.Run(ctx, args, stdin)
}
func lifecycleTransaction(statePath string) string {
value := fmt.Sprintf("%s\x00%d\x00%d", filepath.Clean(statePath), os.Getpid(), time.Now().UnixNano())
sum := sha256.Sum256([]byte(value))
return fmt.Sprintf("%x", sum[:8])
}
func lifecycleImageTag(transaction, role string) string {
return "thothii-core:thothctl-" + transaction + "-" + role
}
func lifecycleOverridePath(statePath, transaction string) string {
if transaction == "" {
transaction = "recovery"
}
return filepath.Join(filepath.Dir(statePath), "pi-lifecycle-"+transaction+".yaml")
}
func writeLifecycleOverride(path, image string) error {
quoted, err := json.Marshal(image)
if err != nil {
return errors.New("lifecycle image override could not be encoded")
}
contents := []byte("services:\n core:\n image: " + string(quoted) + "\n")
if err := writeFileDurably(path, ".pi-lifecycle-", contents); err != nil {
return errors.New("lifecycle image override could not be written durably")
}
return nil
}
// RecoverMaintenance clears a stale durable gate only after the running core and terminal
// recovery metadata prove that no rollback is still required.
func RecoverMaintenance(ctx context.Context, runner Runner, statePath string, confirm bool) error {
if !confirm {
return ErrConfirmationRequired
}
lock, err := acquireLock(statePath)
if err != nil {
return err
}
defer lock.Release()
state, stateErr := readState(statePath)
if stateErr == nil {
if state.Phase != PhaseVerified && state.Phase != PhaseRolledBack && state.Phase != PhaseNoop {
return ErrInterruptedUpdate
}
if err := durableRemove(lifecycleOverridePath(statePath, state.Transaction)); err != nil {
return errors.New("maintenance recovery could not remove the lifecycle override")
}
} else if !errors.Is(stateErr, os.ErrNotExist) {
return stateErr
}
status, err := MaintenanceStatus(ctx, runner)
if err != nil {
return err
}
if !status.Active {
return nil
}
if err := Doctor(ctx, runner); err != nil {
return err
}
return setMaintenance(ctx, runner, false)
}
func sameStrings(left, right []string) bool {
left, right = append([]string(nil), left...), append([]string(nil), right...)
sort.Strings(left)
+425 -37
View File
@@ -7,12 +7,29 @@ import (
"io"
"os"
"path/filepath"
"strconv"
"strings"
"testing"
"github.com/aritmolab/thothii/tools/thothctl/internal/compose"
)
func TestActiveSessionsParsesAuthenticatedBackendBareArrayFixture(t *testing.T) {
contents, err := os.ReadFile(filepath.Join("..", "..", "..", "..", "backend", "test", "fixtures", "sessions-scope-all.json"))
if err != nil {
t.Fatal(err)
}
fake := newFakeRunner()
fake.sessionsWire = string(contents)
active, err := activeSessions(context.Background(), fake)
if err != nil {
t.Fatalf("activeSessions() error = %v", err)
}
if !active {
t.Fatal("activeSessions() = false, want open session from backend wire fixture")
}
}
func TestUpdateBuildsPinnedVersionRecreatesOnlyCoreAndPersistsRecoveryState(t *testing.T) {
fake := newFakeRunner()
dir := t.TempDir()
@@ -28,8 +45,8 @@ func TestUpdateBuildsPinnedVersionRecreatesOnlyCoreAndPersistsRecoveryState(t *t
if result.Phase != PhaseVerified {
t.Fatalf("phase = %q, want %q", result.Phase, PhaseVerified)
}
assertCalled(t, fake.calls, "compose build --pull --build-arg PI_VERSION=0.81.0 core")
assertCalled(t, fake.calls, "compose up --detach --wait --wait-timeout 45 --no-deps --force-recreate core")
assertCalled(t, fake.calls, "build --pull --build-arg PI_VERSION=0.81.0 core")
assertCalled(t, fake.calls, "up --detach --wait --wait-timeout 45 --no-deps --force-recreate core")
assertNotCalled(t, fake.calls, "frontend")
if got := string(readStateBytes(t, result.StatePath)); strings.Contains(got, "secret") || !strings.Contains(got, `"phase": "verified"`) {
t.Fatalf("state = %q, want credential-free verified metadata", got)
@@ -39,15 +56,93 @@ func TestUpdateBuildsPinnedVersionRecreatesOnlyCoreAndPersistsRecoveryState(t *t
}
}
func TestUpdateUsesATransactionScopedComposeOverrideWithoutMutatingTheConfiguredImage(t *testing.T) {
fake := newFakeRunner()
statePath := filepath.Join(t.TempDir(), ".thothctl", "update-state.json")
if _, err := Update(context.Background(), fake, Request{StatePath: statePath, Version: "0.81.0", Source: BuildSource, Confirm: true}); err != nil {
t.Fatal(err)
}
if fake.buildReference == "" || fake.buildReference == fake.configuredImage || !strings.Contains(fake.buildReference, "thothctl-") {
t.Fatalf("build reference = %q, configured = %q; want unique lifecycle tag", fake.buildReference, fake.configuredImage)
}
assertNotCalled(t, fake.calls, "image tag sha256:old "+fake.configuredImage)
if matches, err := filepath.Glob(filepath.Join(filepath.Dir(statePath), "pi-lifecycle-*.yaml")); err != nil || len(matches) != 0 {
t.Fatalf("terminal lifecycle overrides = %v, error = %v; want none", matches, err)
}
}
func TestTwoInstallationsSharingAConfiguredTagUseDifferentLifecycleTags(t *testing.T) {
first, second := newFakeRunner(), newFakeRunner()
for _, item := range []struct {
fake *fakeRunner
path string
}{
{first, filepath.Join(t.TempDir(), "one", "state.json")},
{second, filepath.Join(t.TempDir(), "two", "state.json")},
} {
if _, err := Update(context.Background(), item.fake, Request{StatePath: item.path, Version: "0.81.0", Source: BuildSource, Confirm: true}); err != nil {
t.Fatal(err)
}
}
if first.buildReference == second.buildReference {
t.Fatalf("installations reused lifecycle tag %q", first.buildReference)
}
}
func TestDigestPinnedConfiguredImageIsNeverUsedAsARollbackTagTarget(t *testing.T) {
fake := newFakeRunner()
fake.configuredImage = "registry.example.invalid/core@sha256:" + strings.Repeat("b", 64)
fake.tags = map[string]string{fake.configuredImage: "sha256:old"}
fake.fail = "health"
_, _ = Update(context.Background(), fake, Request{StatePath: filepath.Join(t.TempDir(), "state.json"), Version: "0.81.0", Source: BuildSource, Confirm: true})
assertNotCalled(t, fake.calls, "image tag sha256:old "+fake.configuredImage)
}
func TestMaintenanceLostResponsesAreResolvedByStatusAndEveryRecreateStartsGated(t *testing.T) {
for _, lost := range []string{"activate", "deactivate"} {
t.Run(lost, func(t *testing.T) {
fake := newFakeRunner()
fake.lostMaintenanceResponse = lost
if _, err := Update(context.Background(), fake, Request{StatePath: filepath.Join(t.TempDir(), "state.json"), Version: "0.81.0", Source: BuildSource, Confirm: true}); err != nil {
t.Fatal(err)
}
assertCalled(t, fake.calls, "/internal/maintenance/status")
for index, active := range fake.maintenanceAtRecreate {
if !active {
t.Fatalf("recreate %d started without durable maintenance", index+1)
}
}
})
}
}
func TestCompensationReactivatesMaintenanceAndRescansBeforeRollback(t *testing.T) {
fake := newFakeRunner()
fake.fail = "version"
fake.dropMaintenanceAfterCandidate = true
_, _ = Update(context.Background(), fake, Request{StatePath: filepath.Join(t.TempDir(), "state.json"), Version: "0.81.0", Source: BuildSource, Confirm: true})
rollback := lastCallIndexBefore(fake.calls, "image tag sha256:old", len(fake.calls))
if rollback < 0 {
t.Fatalf("calls %v contain no rollback", fake.calls)
}
recreate := callIndex(fake.calls, "force-recreate core")
activate := lastCallIndexBefore(fake.calls, "/internal/maintenance/activate", rollback)
scan := lastCallIndexBefore(fake.calls, "/sessions?scope=all", rollback)
if activate <= recreate || scan <= recreate {
t.Fatalf("calls %v do not reactivate/confirm and rescan after candidate recreate before rollback", fake.calls)
}
}
func TestUpdatePullsOnlyDigestPinnedSource(t *testing.T) {
fake := newFakeRunner()
digest := "registry.example.invalid/thothii-core@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"
_, err := Update(context.Background(), fake, Request{StatePath: filepath.Join(t.TempDir(), "state.json"), Version: "0.81.0", Source: PullSource, Image: digest, Confirm: true})
if err == nil {
t.Fatal("Update() error = nil, want Pi version verification failure from unchanged fake image")
if err != nil {
t.Fatalf("Update() pull error = %v", err)
}
assertCalled(t, fake.calls, "pull "+digest)
assertCalled(t, fake.calls, "image tag "+digest+" thothii-core:local")
assertCalled(t, fake.calls, "image tag "+digest+" thothii-core:thothctl-")
assertNotCalled(t, fake.calls, "image tag "+digest+" thothii-core:local")
fake = newFakeRunner()
_, err = Update(context.Background(), fake, Request{StatePath: filepath.Join(t.TempDir(), "state.json"), Version: "0.81.0", Source: PullSource, Image: "registry.example.invalid/thothii-core:latest", Confirm: true})
@@ -71,7 +166,7 @@ func TestUpdateIsNoOpWhenDesiredVersionAlreadyRuns(t *testing.T) {
}
func TestUpdateRollsBackAfterPostRecreateFailures(t *testing.T) {
for _, failure := range []string{"health", "version", "smoke"} {
for _, failure := range []string{"recreate", "health", "version", "smoke", "config-drift", "mount-drift"} {
t.Run(failure, func(t *testing.T) {
fake := newFakeRunner()
fake.fail = failure
@@ -83,11 +178,9 @@ func TestUpdateRollsBackAfterPostRecreateFailures(t *testing.T) {
if result.Phase != PhaseRolledBack {
t.Fatalf("phase = %q, want %q", result.Phase, PhaseRolledBack)
}
assertCalled(t, fake.calls, "image tag sha256:old thothii-core:local")
assertCalled(t, fake.calls, "compose up --detach --wait --wait-timeout 45 --no-deps --force-recreate core")
if strings.Join(fake.volumes, ",") != "settings,pi-state,sessions,workspace-registry" {
t.Fatalf("volumes changed: %v", fake.volumes)
}
assertCalled(t, fake.calls, "image tag sha256:old thothii-core:thothctl-")
assertNotCalled(t, fake.calls, "image tag sha256:old thothii-core:local")
assertCalled(t, fake.calls, "up --detach --wait --wait-timeout 45 --no-deps --force-recreate core")
if got := string(readStateBytes(t, statePath)); !strings.Contains(got, `"phase": "rolled_back"`) {
t.Fatalf("state = %q, want rollback metadata", got)
}
@@ -95,6 +188,143 @@ func TestUpdateRollsBackAfterPostRecreateFailures(t *testing.T) {
}
}
func TestEveryRecoveryStateWriteFailureIsHandledTransactionally(t *testing.T) {
for failAt := 1; failAt <= 4; failAt++ {
t.Run(fmt.Sprintf("write-%d", failAt), func(t *testing.T) {
fake := newFakeRunner()
writes := 0
hooks := defaultLifecycleHooks
hooks.writeState = func(path string, state State) error {
writes++
if writes == failAt {
return errors.New("injected state write failure")
}
return writeState(path, state)
}
result, err := updateWithHooks(context.Background(), fake, Request{
StatePath: filepath.Join(t.TempDir(), "state.json"), Version: "0.81.0", Source: BuildSource, Confirm: true,
}, hooks)
if err == nil {
t.Fatal("updateWithHooks() error = nil, want injected state failure")
}
if fake.currentImage != "sha256:old" {
t.Fatalf("current image = %q, want restored previous", fake.currentImage)
}
if failAt > 1 && result.Phase != PhaseRolledBack {
t.Fatalf("phase = %q, want rolled_back", result.Phase)
}
if fake.maintenance {
t.Fatal("maintenance remained active after proven stable recovery")
}
})
}
}
func TestCompensationWriteFailureKeepsMaintenanceActiveForExplicitRecovery(t *testing.T) {
fake := newFakeRunner()
fake.fail = "health"
writes := 0
hooks := defaultLifecycleHooks
hooks.writeState = func(path string, state State) error {
writes++
if writes == 4 {
return errors.New("injected compensation state write failure")
}
return writeState(path, state)
}
result, err := updateWithHooks(context.Background(), fake, Request{
StatePath: filepath.Join(t.TempDir(), "state.json"), Version: "0.81.0", Source: BuildSource, Confirm: true,
}, hooks)
if err == nil || result.Phase != PhaseFailed {
t.Fatalf("result=%+v error=%v, want failed recovery", result, err)
}
if !fake.maintenance {
t.Fatal("maintenance was cleared without durable rollback state")
}
}
func TestMaintenanceClearAndCompensationFailuresRemainGated(t *testing.T) {
for _, failure := range []string{"maintenance-clear", "compensation"} {
t.Run(failure, func(t *testing.T) {
fake := newFakeRunner()
fake.fail = failure
result, err := Update(context.Background(), fake, Request{StatePath: filepath.Join(t.TempDir(), "state.json"), Version: "0.81.0", Source: BuildSource, Confirm: true})
if err == nil || result.Phase != PhaseFailed {
t.Fatalf("result=%+v error=%v", result, err)
}
if !fake.maintenance {
t.Fatal("maintenance was cleared after an unverified terminal failure")
}
})
}
}
func TestRecoverMaintenanceClearsOnlyAfterTerminalStateAndVerifiedSmoke(t *testing.T) {
fake := newFakeRunner()
fake.maintenance = true
statePath := filepath.Join(t.TempDir(), "state.json")
previous := stateImageForTest(t, fake)
state := State{Transaction: "recover-test", Phase: PhaseVerified, Previous: previous}
writeStateForTest(t, statePath, state)
overridePath := lifecycleOverridePath(statePath, state.Transaction)
if err := writeLifecycleOverride(overridePath, previous.Reference); err != nil {
t.Fatal(err)
}
if err := RecoverMaintenance(context.Background(), fake, statePath, true); err != nil {
t.Fatalf("RecoverMaintenance() error = %v", err)
}
if fake.maintenance {
t.Fatal("maintenance remained active after verified terminal recovery")
}
if _, err := os.Stat(overridePath); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("lifecycle override still exists: %v", err)
}
assertCalled(t, fake.calls, "/models")
assertCalled(t, fake.calls, "/settings")
}
func TestRecoverMaintenanceRefusesPendingTransaction(t *testing.T) {
fake := newFakeRunner()
fake.maintenance = true
statePath := filepath.Join(t.TempDir(), "state.json")
writeStateForTest(t, statePath, State{Phase: PhaseRecreated, Previous: stateImageForTest(t, fake)})
err := RecoverMaintenance(context.Background(), fake, statePath, true)
if !errors.Is(err, ErrInterruptedUpdate) {
t.Fatalf("RecoverMaintenance() error = %v, want ErrInterruptedUpdate", err)
}
if !fake.maintenance {
t.Fatal("pending transaction maintenance was cleared")
}
}
func TestRollbackFinalStateWriteFailureKeepsMaintenanceAndOverrideForRecovery(t *testing.T) {
fake := newFakeRunner()
statePath := filepath.Join(t.TempDir(), "state.json")
previous := stateImageForTest(t, fake)
previous.Reference = "thothii-core:thothctl-rollback-test-previous"
fake.tags[previous.Reference] = previous.ID
writeStateForTest(t, statePath, State{
Transaction: "rollback-test",
Phase: PhaseRecreated,
Previous: previous,
})
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)
if err == nil || result.Phase != PhaseFailed {
t.Fatalf("rollbackWithHooks() = %+v, %v; want failed durable finalization", result, err)
}
if !fake.maintenance {
t.Fatal("maintenance was cleared without durable rollback finalization")
}
if _, err := os.Stat(lifecycleOverridePath(statePath, "rollback-test")); err != nil {
t.Fatalf("recovery override was not preserved: %v", err)
}
}
func TestUpdateDoesNotRecreateWhenPreflightOrBuildFails(t *testing.T) {
for _, failure := range []string{"preflight", "build"} {
t.Run(failure, func(t *testing.T) {
@@ -104,9 +334,15 @@ func TestUpdateDoesNotRecreateWhenPreflightOrBuildFails(t *testing.T) {
if err == nil {
t.Fatal("Update() error = nil, want failure")
}
if failure == "preflight" && result.Phase == PhaseRolledBack { t.Fatalf("preflight failure unexpectedly rolled back: %+v", result) }
if failure == "build" && result.Phase != PhaseRolledBack { t.Fatalf("candidate build failure must compensate: %+v", result) }
if failure == "preflight" { assertNotCalled(t, fake.calls, "force-recreate") }
if failure == "preflight" && result.Phase == PhaseRolledBack {
t.Fatalf("preflight failure unexpectedly rolled back: %+v", result)
}
if failure == "build" && result.Phase != PhaseRolledBack {
t.Fatalf("candidate build failure must compensate: %+v", result)
}
if failure == "preflight" {
assertNotCalled(t, fake.calls, "force-recreate")
}
})
}
}
@@ -149,6 +385,8 @@ func TestRollbackRestoresInterruptedOrPreviouslyRecordedState(t *testing.T) {
t.Fatal(err)
}
previous.ConfigurationSHA = configured.ConfigurationSHA
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)
if err != nil {
@@ -157,14 +395,14 @@ func TestRollbackRestoresInterruptedOrPreviouslyRecordedState(t *testing.T) {
if result.Phase != PhaseRolledBack {
t.Fatalf("phase = %q, want %q", result.Phase, PhaseRolledBack)
}
assertCalled(t, fake.calls, "image tag sha256:old thothii-core:local")
assertCalled(t, fake.calls, "compose up --detach --wait --wait-timeout 45 --no-deps --force-recreate core")
assertCalled(t, fake.calls, "image tag sha256:old thothii-core:thothctl-test-previous")
assertCalled(t, fake.calls, "up --detach --wait --wait-timeout 45 --no-deps --force-recreate core")
}
func TestUpdateRefusesToOverwriteInterruptedRecoveryState(t *testing.T) {
fake := newFakeRunner()
statePath := filepath.Join(t.TempDir(), "state.json")
writeStateForTest(t, statePath, State{Phase: PhaseRecreated, Previous: Image{ID: "sha256:old", Reference: "thothii-core:local", Volumes: []string{"settings"}, MountFingerprint: mountFingerprint(nil)}})
writeStateForTest(t, statePath, State{Phase: PhaseRecreated, Previous: Image{ID: "sha256:old", Reference: "thothii-core:local", MountFingerprint: mountFingerprint(nil)}})
_, err := Update(context.Background(), fake, Request{StatePath: statePath, Version: "0.81.0", Source: BuildSource, Confirm: true})
if !errors.Is(err, ErrInterruptedUpdate) {
t.Fatalf("Update() error = %v, want interrupted update error", err)
@@ -189,43 +427,63 @@ func TestRunningImageCapturesServerBindAndNamedMountIdentity(t *testing.T) {
func TestCanonicalDigestReferenceRejectsCredentialsAndURLForms(t *testing.T) {
valid := "registry.example.invalid/thothii-core@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"
if got, err := canonicalDigestReference(valid); err != nil || got != valid { t.Fatalf("canonicalDigestReference() = %q, %v", got, err) }
if got, err := canonicalDigestReference(valid); err != nil || got != valid {
t.Fatalf("canonicalDigestReference() = %q, %v", got, err)
}
for _, invalid := range []string{"https://registry.example.invalid/a@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", "user:pass@registry.example/a@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", "registry.example/a@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa?token=x"} {
if _, err := canonicalDigestReference(invalid); err == nil { t.Fatalf("accepted unsafe reference %q", invalid) }
if _, err := canonicalDigestReference(invalid); err == nil {
t.Fatalf("accepted unsafe reference %q", invalid)
}
}
}
type fakeRunner struct {
calls []string
fail string
version string
activeSessions bool
built bool
currentImage string
volumes []string
mountsJSON string
calls []string
fail string
version string
activeSessions bool
built bool
currentImage string
mountsJSON string
sessionsWire string
configuredImage string
buildReference string
tags map[string]string
imageVersions map[string]string
maintenance bool
maintenanceAtRecreate []bool
lostMaintenanceResponse string
dropMaintenanceAfterCandidate bool
modelsWire string
rollbackPrepared bool
}
func newFakeRunner() *fakeRunner {
return &fakeRunner{version: "0.80.3", currentImage: "sha256:old", volumes: []string{"settings", "pi-state", "sessions", "workspace-registry"}}
return &fakeRunner{
version: "0.80.3", currentImage: "sha256:old", configuredImage: "thothii-core:local",
tags: map[string]string{"thothii-core:local": "sha256:old"},
imageVersions: map[string]string{"sha256:old": "0.80.3"},
}
}
func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose.Result, error) {
call := strings.Join(args, " ")
f.calls = append(f.calls, call)
if strings.Contains(call, "image tag sha256:old") {
if f.built && f.fail != "compensation" && strings.Contains(call, "image tag sha256:old") {
f.fail = ""
f.currentImage = "sha256:old"
}
if f.fail == "preflight" && strings.Contains(call, "config --format json") {
return compose.Result{ExitCode: 1}, errors.New("provider token=secret")
}
if f.fail == "build" && strings.Contains(call, "compose build") {
if f.fail == "build" && containsArg(args, "build") {
return compose.Result{ExitCode: 1}, errors.New("build token=secret")
}
if f.fail == "health" && f.built && strings.Contains(call, "curl -fsS http://127.0.0.1:8787/health") {
return compose.Result{ExitCode: 1}, errors.New("health token=secret")
}
if f.fail == "compensation" && f.built && !f.rollbackPrepared && strings.Contains(call, "curl -fsS http://127.0.0.1:8787/health") {
return compose.Result{ExitCode: 1}, errors.New("candidate health failure")
}
if f.fail == "version" && f.built && strings.Contains(call, "pi --version") && strings.Contains(call, "exec") {
return compose.Result{ExitCode: 1}, errors.New("version token=secret")
}
@@ -234,34 +492,98 @@ func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose
}
switch {
case strings.Contains(call, "config --format json"):
return compose.Result{Stdout: `{"services":{"core":{"image":"thothii-core:local","environment":{"THT_LLM_URL":"https://llm.example.invalid"}}}}`}, nil
endpoint := "https://llm.example.invalid"
if f.fail == "config-drift" && f.currentImage == "sha256:candidate" {
endpoint = "https://drift.example.invalid"
}
return compose.Result{Stdout: `{"services":{"core":{"image":"` + selectedCoreReference(args, f.configuredImage) + `","environment":{"THT_LLM_URL":"` + endpoint + `"}}}}`}, nil
case strings.Contains(call, "ps -q core"):
return compose.Result{Stdout: "core-container\n"}, nil
case strings.Contains(call, "inspect --format {{.Image}}"):
return compose.Result{Stdout: f.currentImage + "\n"}, nil
case strings.Contains(call, "inspect --format {{json .Mounts}}"):
if f.fail == "mount-drift" && f.currentImage == "sha256:candidate" {
return compose.Result{Stdout: `[{"Type":"volume","Name":"wrong-settings","Source":"wrong-settings","Destination":"/data/settings","RW":true}]`}, nil
}
if f.mountsJSON != "" {
return compose.Result{Stdout: f.mountsJSON}, nil
}
return compose.Result{Stdout: `[{"Type":"volume","Name":"settings","Source":"settings","Destination":"/data/settings","RW":true},{"Type":"volume","Name":"pi-state","Source":"pi-state","Destination":"/home/thoth/.pi","RW":true},{"Type":"volume","Name":"sessions","Source":"sessions","Destination":"/data/sessions","RW":true},{"Type":"volume","Name":"workspace-registry","Source":"workspace-registry","Destination":"/data/workspace-registry","RW":true}]`}, nil
case strings.Contains(call, "/internal/maintenance/activate"):
f.maintenance = true
if f.lostMaintenanceResponse == "activate" {
f.lostMaintenanceResponse = ""
return compose.Result{ExitCode: 52}, errors.New("lost activation response")
}
return compose.Result{Stdout: `{"active":true,"admissions":0}`}, nil
case strings.Contains(call, "/internal/maintenance/deactivate"):
if f.fail == "maintenance-clear" {
return compose.Result{ExitCode: 53}, errors.New("maintenance clear failure")
}
f.maintenance = false
if f.lostMaintenanceResponse == "deactivate" {
f.lostMaintenanceResponse = ""
return compose.Result{ExitCode: 52}, errors.New("lost deactivation response")
}
return compose.Result{Stdout: `{"active":false,"admissions":0}`}, nil
case strings.Contains(call, "/internal/maintenance/status"):
return compose.Result{Stdout: fmt.Sprintf(`{"active":%t,"admissions":0}`, f.maintenance)}, nil
case strings.Contains(call, "/sessions?scope=all"):
if f.sessionsWire != "" {
return compose.Result{Stdout: f.sessionsWire}, nil
}
if f.activeSessions {
f.activeSessions = false
return compose.Result{Stdout: `{"sessions":[{"status":"open","archived":false}]}`}, nil
return compose.Result{Stdout: `[{"status":"open","archived":false}]`}, nil
}
return compose.Result{Stdout: `{"sessions":[]}`}, nil
case strings.Contains(call, "compose build"):
return compose.Result{Stdout: `[]`}, nil
case containsArg(args, "build"):
f.built = true
f.version = "0.81.0"
f.currentImage = "sha256:candidate"
f.buildReference = selectedCoreReference(args, f.configuredImage)
f.tags[f.buildReference] = "sha256:candidate"
f.imageVersions["sha256:candidate"] = "0.81.0"
return compose.Result{}, nil
case len(args) == 2 && args[0] == "pull":
f.tags[args[1]] = "sha256:candidate"
f.imageVersions["sha256:candidate"] = "0.81.0"
return compose.Result{}, nil
case len(args) >= 4 && args[0] == "image" && args[1] == "tag":
source, target := args[2], args[3]
id := source
if tagged, ok := f.tags[source]; ok {
id = tagged
}
f.tags[target] = id
if f.built && id == "sha256:old" {
f.rollbackPrepared = true
}
return compose.Result{}, nil
case containsArg(args, "up"):
f.maintenanceAtRecreate = append(f.maintenanceAtRecreate, f.maintenance)
reference := selectedCoreReference(args, f.configuredImage)
if id, ok := f.tags[reference]; ok {
f.currentImage = id
}
if version, ok := f.imageVersions[f.currentImage]; ok {
f.version = version
}
if f.dropMaintenanceAfterCandidate && f.currentImage == "sha256:candidate" {
f.maintenance = false
f.dropMaintenanceAfterCandidate = false
}
if f.fail == "recreate" && f.currentImage == "sha256:candidate" {
return compose.Result{ExitCode: 54}, errors.New("recreate failure")
}
if f.fail == "compensation" && f.rollbackPrepared {
return compose.Result{ExitCode: 55}, errors.New("rollback recreate failure")
}
return compose.Result{}, nil
case strings.Contains(call, "pi --version"):
return compose.Result{Stdout: f.version + "\n"}, nil
case strings.Contains(call, "/models"):
if f.modelsWire != "" {
return compose.Result{Stdout: f.modelsWire}, nil
}
return compose.Result{Stdout: `{"models":[{"id":"model","provider":"provider"}]}`}, nil
case strings.Contains(call, "/settings"):
return compose.Result{Stdout: `{"provider":"provider","model":"model","thinking":"medium"}`}, nil
@@ -271,6 +593,57 @@ func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose
return compose.Result{}, nil
}
func containsArg(args []string, wanted string) bool {
for _, arg := range args {
if arg == wanted {
return true
}
}
return false
}
func selectedCoreReference(args []string, fallback string) string {
for index := 0; index+1 < len(args); index++ {
if args[index] != "-f" || !strings.Contains(filepath.Base(args[index+1]), "pi-lifecycle-") {
continue
}
contents, err := os.ReadFile(args[index+1])
if err != nil {
continue
}
for _, line := range strings.Split(string(contents), "\n") {
line = strings.TrimSpace(line)
if !strings.HasPrefix(line, "image:") {
continue
}
value := strings.TrimSpace(strings.TrimPrefix(line, "image:"))
if decoded, err := strconv.Unquote(value); err == nil {
return decoded
}
return value
}
}
return fallback
}
func callIndex(calls []string, contains string) int {
for index, call := range calls {
if strings.Contains(call, contains) {
return index
}
}
return -1
}
func lastCallIndexBefore(calls []string, contains string, before int) int {
for index := before - 1; index >= 0; index-- {
if strings.Contains(calls[index], contains) {
return index
}
}
return -1
}
func assertCalled(t *testing.T, calls []string, want string) {
t.Helper()
for _, call := range calls {
@@ -302,3 +675,18 @@ func writeStateForTest(t *testing.T, path string, state State) {
t.Fatal(err)
}
}
func stateImageForTest(t *testing.T, fake *fakeRunner) Image {
t.Helper()
configured, err := renderedCore(context.Background(), fake)
if err != nil {
t.Fatal(err)
}
image, err := runningImage(context.Background(), fake, fake.configuredImage)
if err != nil {
t.Fatal(err)
}
image.ConfigurationSHA = configured.ConfigurationSHA
image.Reference = "thothii-core:thothctl-test-previous"
return image
}