fix: harden pi maintenance lifecycle

This commit is contained in:
2026-08-04 19:09:04 +02:00
parent 8fde1f81c7
commit 0b9ad7f53f
21 changed files with 432 additions and 171 deletions
+14 -14
View File
@@ -23,7 +23,7 @@ const usage = `Usage: thothctl --installation <absolute-path>/thothii-installati
Commands:
status Show the Compose service state.
doctor Validate Docker, Compose, rendered configuration, line endings, volumes, and health.
logs [--follow] Show sanitized service logs (the default is the latest 200 lines).
logs Show the latest 200 sanitized service log lines.
start Start the installation in the background.
stop Stop the installation.
update --check-only Validate the current installation without changing containers.
@@ -31,8 +31,8 @@ Commands:
pi doctor Check Pi preconditions without changing the installation.
pi test Run the temporary Pi/core smoke checks.
pi check Alias for pi test.
pi configure Store non-secret Pi defaults (credentials stay in PI_AUTH_FILE).
pi update Rebuild or pull a pinned Pi image (requires --yes).
pi configure Apply non-secret provider/model/thinking defaults to core (credentials stay in PI_AUTH_FILE).
pi update Rebuild or pull a pinned Pi image (--source build|pull and --yes required).
pi rollback --yes Restore the image recorded by the latest Pi update.
pi logs Show the latest sanitized core logs.
`
@@ -169,10 +169,10 @@ func piCommand(ctx context.Context, installation config.Installation, runner com
if err != nil {
return commandUsageError(stderr, err.Error())
}
if err := pi.Configure(ctx, controlled, filepath.Join(installation.ProjectDirectory, ".thothctl", "pi-defaults.json"), defaults); err != nil {
if err := pi.Configure(ctx, controlled, defaults); err != nil {
return piFailure(stderr, err, secretValues)
}
fmt.Fprintln(stdout, "Pi defaults saved. Put credentials only in the configured PI_AUTH_FILE (mode 0600).")
fmt.Fprintln(stdout, "Pi defaults applied and read back. Put credentials only in PI_AUTH_FILE (/home/thoth/.pi/agent/auth.json, mode 0600); never pass credentials to thothctl.")
return 0
case "update":
request, err := parsePiUpdateArgs(args[1:], filepath.Join(installation.ProjectDirectory, ".thothctl", "update-state.json"))
@@ -219,20 +219,18 @@ func parsePiConfigureArgs(args []string) (pi.Defaults, error) {
value.Model = v
case "--thinking":
value.Thinking = v
case "--llm-url":
value.LLMURL = v
default:
return pi.Defaults{}, fmt.Errorf("unknown pi configure option %q", key)
}
}
if value.Provider == "" || value.Model == "" || value.Thinking == "" || value.LLMURL == "" {
return pi.Defaults{}, errors.New("pi configure requires --provider --model --thinking --llm-url")
if value.Provider == "" || value.Model == "" || value.Thinking == "" {
return pi.Defaults{}, errors.New("pi configure requires --provider --model --thinking; THT_LLM_URL stays Compose-managed")
}
return value, nil
}
func parsePiUpdateArgs(args []string, statePath string) (pi.Request, error) {
request := pi.Request{StatePath: statePath, Source: pi.BuildSource}
request := pi.Request{StatePath: statePath}
for len(args) > 0 {
switch args[0] {
case "--version":
@@ -264,15 +262,17 @@ func parsePiUpdateArgs(args []string, statePath string) (pi.Request, error) {
return pi.Request{}, fmt.Errorf("unknown pi update option %q", args[0])
}
}
if request.Version == "" {
return pi.Request{}, errors.New("pi update requires --version <pinned-version>")
}
if request.Version == "" { return pi.Request{}, errors.New("pi update requires --version <pinned-version>") }
if request.Source == "" { return pi.Request{}, errors.New("pi update requires explicit --source build or pull") }
if request.Source != pi.BuildSource && request.Source != pi.PullSource { return pi.Request{}, errors.New("--source requires build or pull") }
if request.Source == pi.PullSource && request.Image == "" { return pi.Request{}, errors.New("--source pull requires --image <digest-reference>") }
if request.Source == pi.BuildSource && request.Image != "" { return pi.Request{}, errors.New("--image is valid only with --source pull") }
return request, nil
}
func piFailure(stderr io.Writer, err error, secretValues []string) int {
code := 1
if errors.Is(err, pi.ErrConfirmationRequired) || errors.Is(err, pi.ErrActiveSessions) || errors.Is(err, pi.ErrInterruptedUpdate) {
if errors.Is(err, pi.ErrConfirmationRequired) || errors.Is(err, pi.ErrInvalidRequest) {
code = 2
}
var childExit interface{ ExitCode() int }
+1 -1
View File
@@ -368,7 +368,7 @@ func TestRunPiUpdateRequiresExplicitConfirmationWithoutInvokingDocker(t *testing
fixture.setEnvironment(t)
var stdout, stderr bytes.Buffer
exitCode := run(context.Background(), []string{"--installation", fixture.installationPath, "pi", "update", "--version", "0.81.0"}, &stdout, &stderr)
exitCode := run(context.Background(), []string{"--installation", fixture.installationPath, "pi", "update", "--version", "0.81.0", "--source", "build"}, &stdout, &stderr)
if exitCode != 2 {
t.Errorf("run() exit code = %d, want 2", exitCode)
+3
View File
@@ -6,6 +6,9 @@ require gopkg.in/yaml.v3 v3.0.1
require (
github.com/compose-spec/compose-go/v2 v2.14.0
github.com/distribution/reference v0.6.0
github.com/sirupsen/logrus v1.9.0
golang.org/x/sys v0.5.0
)
require github.com/opencontainers/go-digest v1.0.0 // indirect
+4
View File
@@ -3,8 +3,12 @@ github.com/compose-spec/compose-go/v2 v2.14.0/go.mod h1:ZU6zlcweCZKyiB7BVfCizQT9
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/distribution/reference v0.6.0 h1:0IXCQ5g4/QMHHkarYzh5l+u8T3t73zM5QvfrDyIgxBk=
github.com/distribution/reference v0.6.0/go.mod h1:BbU0aIcezP1/5jX/8MP0YiH4SdvB5Y4f/wlDRiLyi3E=
github.com/google/go-cmp v0.5.9 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38=
github.com/google/go-cmp v0.5.9/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8Oi/yOhh5U=
github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/sirupsen/logrus v1.9.0 h1:trlNQbNUG3OdDrDil03MCb1H2o9nJ1x4/5LYw7byDE0=
+48 -46
View File
@@ -7,9 +7,6 @@ import (
"errors"
"fmt"
"io"
"net/url"
"os"
"path/filepath"
"regexp"
"strings"
@@ -22,22 +19,29 @@ type Defaults struct {
Provider string `json:"provider"`
Model string `json:"model"`
Thinking string `json:"thinking"`
LLMURL string `json:"llm_url"`
}
// Configure validates and atomically stores only local non-secret Pi defaults.
func Configure(ctx context.Context, runner Runner, path string, value Defaults) error {
var internalIdentityHeaders = []string{
"-H", "x-thoth-principal-issuer: thothctl",
"-H", "x-thoth-principal-subject: thothctl-maintenance",
"-H", "x-thoth-principal-display-name: Thothctl maintenance",
"-H", "x-thoth-is-admin: 1",
}
// 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) error {
if !choicePattern.MatchString(value.Provider) || !choicePattern.MatchString(value.Model) {
return errors.New("provider and model must be supported identifiers")
}
if value.Thinking != "low" && value.Thinking != "medium" && value.Thinking != "high" {
return errors.New("thinking must be low, medium, or high")
}
u, err := url.Parse(value.LLMURL)
if err != nil || (u.Scheme != "https" && u.Scheme != "http") || u.Host == "" || u.User != nil || u.RawQuery != "" || u.Fragment != "" {
return errors.New("LLM endpoint must be an http(s) URL without credentials, query, or fragment")
}
models, err := runCompose(ctx, runner, "exec", "-T", "core", "curl", "-fsS", "http://127.0.0.1:8787/models")
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)
}
@@ -59,39 +63,19 @@ func Configure(ctx context.Context, runner Runner, path string, value Defaults)
if !found {
return errors.New("provider/model is not in Pi options")
}
return writeJSON(path, value)
}
func writeJSON(path string, value any) error {
contents, err := json.MarshalIndent(value, "", " ")
if err != nil {
return err
}
contents = append(contents, '\n')
if err = os.MkdirAll(filepath.Dir(path), 0o700); err != nil {
return errors.New("could not create Pi configuration directory")
}
temporary, err := os.CreateTemp(filepath.Dir(path), ".pi-defaults-*.tmp")
if err != nil {
return errors.New("could not write Pi configuration")
}
name := temporary.Name()
defer os.Remove(name)
if err = temporary.Chmod(0o600); err == nil {
_, err = temporary.Write(contents)
}
if err == nil {
err = temporary.Sync()
}
if closeErr := temporary.Close(); err == nil {
err = closeErr
}
if err == nil {
err = os.Rename(name, path)
}
if err != nil {
return errors.New("could not atomically write Pi configuration")
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) }
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 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")
}
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") }
return nil
}
@@ -131,7 +115,7 @@ func Doctor(ctx context.Context, runner Runner) error {
return commandError("Pi preflight check", result, err)
}
}
return nil
return Test(ctx, runner)
}
// Test performs the pre-Task-8 composite smoke through core's private loopback endpoint.
@@ -140,7 +124,10 @@ func Test(ctx context.Context, runner Runner) error {
return err
}
for _, path := range []string{"health", "models", "settings"} {
result, err := runCompose(ctx, runner, "exec", "-T", "core", "curl", "-fsS", "http://127.0.0.1:8787/"+path)
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)
}
@@ -148,13 +135,28 @@ func Test(ctx context.Context, runner Runner) error {
if err := json.Unmarshal([]byte(result.Stdout), &payload); err != nil {
return fmt.Errorf("Pi smoke check returned invalid %s response", path)
}
if _, ok := payload.(map[string]any); !ok {
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") }
}
}
return nil
}
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 {
+7 -7
View File
@@ -2,7 +2,6 @@ package pi
import (
"context"
"path/filepath"
"strings"
"testing"
)
@@ -17,16 +16,17 @@ func TestDoctorRequiresExternalEndpointAuthPiStateAndHealth(t *testing.T) {
}
}
func TestConfigureValidatesBackendModelOptionsAndWritesNoSecrets(t *testing.T) {
func TestConfigureValidatesBackendModelOptionsWritesRealCoreSettingsAndUsesUpstreamIdentity(t *testing.T) {
fake := newFakeRunner()
path := filepath.Join(t.TempDir(), "pi-defaults.json")
if err := Configure(context.Background(), fake, path, Defaults{Provider: "provider", Model: "model", Thinking: "medium", LLMURL: "https://llm.example.invalid"}); err != nil {
if err := Configure(context.Background(), fake, Defaults{Provider: "provider", Model: "model", Thinking: "medium"}); err != nil {
t.Fatal(err)
}
if got := string(readStateBytes(t, path)); strings.Contains(got, "secret") || !strings.Contains(got, "llm.example.invalid") {
t.Fatalf("defaults=%q", got)
assertCalled(t, fake.calls, "node /app/backend/dist/settings/settings-cli.js --provider provider --model model --thinking medium")
assertCalled(t, fake.calls, "x-thoth-principal-subject: thothctl-maintenance")
if got := strings.Join(fake.calls, "\n"); strings.Contains(got, "pi-defaults.json") || strings.Contains(got, "secret") {
t.Fatalf("commands=%q", got)
}
if err := Configure(context.Background(), fake, path, Defaults{Provider: "provider", Model: "unknown", Thinking: "medium", LLMURL: "https://llm.example.invalid"}); err == nil {
if err := Configure(context.Background(), fake, Defaults{Provider: "provider", Model: "unknown", Thinking: "medium"}); err == nil {
t.Fatal("expected unknown model rejection")
}
}
@@ -0,0 +1,15 @@
//go:build !windows
package pi
import "os"
// 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 }
dir, err := os.Open(directory)
if err != nil { return err }
defer dir.Close()
return dir.Sync()
}
@@ -0,0 +1,15 @@
//go:build windows
package pi
import "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 }
to, err := windows.UTF16PtrFromString(target)
if err != nil { return err }
return windows.MoveFileEx(from, to, windows.MOVEFILE_REPLACE_EXISTING|windows.MOVEFILE_WRITE_THROUGH)
}
@@ -0,0 +1,16 @@
//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)
}
@@ -0,0 +1,14 @@
//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
}
+77 -40
View File
@@ -5,13 +5,15 @@ import (
"encoding/json"
"errors"
"fmt"
"crypto/sha256"
"os"
"path/filepath"
"runtime"
"sort"
"strings"
"time"
)
const stateFileVersion = 1
const stateFileVersion = 2
// Phase describes the durable point reached by a Pi update.
type Phase string
@@ -32,6 +34,7 @@ type Image struct {
Reference string `json:"reference"`
Volumes []string `json:"volumes"`
Mounts []Mount `json:"mounts"`
MountFingerprint string `json:"mount_fingerprint"`
ConfigurationSHA string `json:"configuration_sha256,omitempty"`
}
@@ -39,9 +42,10 @@ type Image struct {
type Mount struct {
Type string `json:"type"`
Name string `json:"name,omitempty"`
Source string `json:"source"`
SourceSHA256 string `json:"source_sha256"`
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
@@ -72,14 +76,14 @@ func readState(path string) (State, error) {
if err := json.Unmarshal(contents, &state); err != nil {
return State{}, errors.New("update recovery state is invalid")
}
if state.Version != stateFileVersion || state.Previous.ID == "" || state.Previous.Reference == "" {
if state.Version != stateFileVersion || state.Previous.ID == "" || state.Previous.Reference == "" || state.Previous.MountFingerprint == "" {
return State{}, errors.New("update recovery state is incomplete")
}
return state, nil
}
func writeState(path string, state State) error {
if state.Previous.ID == "" || state.Previous.Reference == "" {
if state.Previous.ID == "" || state.Previous.Reference == "" || state.Previous.MountFingerprint == "" {
return errors.New("refusing to write incomplete update recovery state")
}
state.Version = stateFileVersion
@@ -89,45 +93,52 @@ func writeState(path string, state State) error {
return fmt.Errorf("encode update recovery state: %w", err)
}
contents = append(contents, '\n')
directory := filepath.Dir(path)
if err := os.MkdirAll(directory, 0o700); err != nil {
return errors.New("could not create update recovery directory")
}
temporary, err := os.CreateTemp(directory, ".update-state-*.tmp")
if err != nil {
return errors.New("could not write update recovery state")
}
temporaryName := temporary.Name()
defer os.Remove(temporaryName)
if err := temporary.Chmod(0o600); err != nil {
temporary.Close()
return errors.New("could not protect update recovery state")
}
if _, err := temporary.Write(contents); err != nil {
temporary.Close()
return errors.New("could not write update recovery state")
}
if err := temporary.Sync(); err != nil {
temporary.Close()
return errors.New("could not durably write update recovery state")
}
if err := temporary.Close(); err != nil {
return errors.New("could not write update recovery state")
}
if err := os.Rename(temporaryName, path); err != nil {
return errors.New("could not finalize update recovery state")
}
if runtime.GOOS != "windows" {
if directoryHandle, err := os.Open(directory); err == nil {
_ = directoryHandle.Sync()
_ = directoryHandle.Close()
}
if err := writeFileDurably(path, ".update-state-", contents); err != nil {
return fmt.Errorf("could not durably write update recovery state: %w", err)
}
return nil
}
func writeFileDurably(path, prefix string, contents []byte) error {
directory := filepath.Dir(path)
if err := os.MkdirAll(directory, 0o700); err != nil { return err }
temporary, err := os.CreateTemp(directory, prefix+"*.tmp")
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 }
return durableReplace(temporaryName, path, directory)
}
func mountSourceHash(source string) string {
sum := sha256.Sum256([]byte(source))
return fmt.Sprintf("%x", sum[:])
}
func mountFingerprint(mounts []Mount) string {
values := make([]string, len(mounts))
for i, mount := range mounts {
values[i] = strings.Join([]string{mount.Type, mount.Name, mount.SourceSHA256, mount.Destination, fmt.Sprint(mount.RW), mount.Options}, "\x00")
}
sort.Strings(values)
sum := sha256.Sum256([]byte(strings.Join(values, "\n")))
return fmt.Sprintf("%x", sum[:])
}
type lockOwner struct {
PID int `json:"pid"`
Host string `json:"host"`
StartedAt time.Time `json:"started_at"`
Transaction string `json:"transaction"`
}
type updateLock struct{ path string }
var ErrLockHeld = errors.New("another Pi update or rollback is already in progress")
func acquireLock(statePath string) (*updateLock, error) {
if err := os.MkdirAll(filepath.Dir(statePath), 0o700); err != nil {
return nil, errors.New("could not create Pi update recovery directory")
@@ -135,10 +146,36 @@ func acquireLock(statePath string) (*updateLock, error) {
path := statePath + ".lock"
if err := os.Mkdir(path, 0o700); err != nil {
if errors.Is(err, os.ErrExist) {
return nil, errors.New("another Pi update or rollback is already in progress; recovery lock retained")
if reclaimDeadLocalLock(path) {
return acquireLock(statePath)
}
return nil, ErrLockHeld
}
return nil, errors.New("could not acquire Pi update lock")
}
host, err := os.Hostname()
if err != nil { _ = os.Remove(path); 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 { _ = os.Remove(path); return nil, errors.New("could not record Pi update lock owner") }
if err := writeFileDurably(filepath.Join(path, "owner.json"), ".owner-", append(contents, '\n')); err != nil {
_ = os.Remove(path)
return nil, errors.New("could not record Pi update lock owner")
}
return &updateLock{path: path}, nil
}
func (l *updateLock) Release() { _ = os.Remove(l.path) }
func (l *updateLock) Release() { _ = os.Remove(filepath.Join(l.path, "owner.json")); _ = 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 string) bool {
contents, err := os.ReadFile(filepath.Join(path, "owner.json"))
if err != nil { return false }
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 }
if err := os.Remove(filepath.Join(path, "owner.json")); err != nil { return false }
return os.Remove(path) == nil
}
+84 -52
View File
@@ -10,14 +10,16 @@ import (
"sort"
"strings"
"time"
"github.com/distribution/reference"
)
var (
ErrConfirmationRequired = errors.New("update requires --yes after reviewing the planned Pi version")
ErrActiveSessions = errors.New("active sessions must be drained before updating Pi; use --drain only after they are complete")
ErrInterruptedUpdate = errors.New("a previous Pi update is incomplete; run pi rollback --yes before starting another update")
ErrInvalidRequest = errors.New("invalid Pi lifecycle request")
versionPattern = regexp.MustCompile(`^[0-9]+(?:\.[0-9]+){1,3}(?:[-+][0-9A-Za-z.-]+)?$`)
digestPattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._/:@-]*@sha256:[a-f0-9]{64}$`)
)
// Source chooses whether the candidate is built from this checkout or pulled from an immutable image.
@@ -45,7 +47,7 @@ type Result struct {
}
// Update performs a recoverable core-only Pi update using the default Compose command layout.
func Update(ctx context.Context, runner Runner, request Request) (Result, error) {
func Update(ctx context.Context, runner Runner, request Request) (result Result, retErr error) {
lock, err := acquireLock(request.StatePath)
if err != nil {
return Result{StatePath: request.StatePath}, err
@@ -58,42 +60,39 @@ func Update(ctx context.Context, runner Runner, request Request) (Result, error)
return Result{StatePath: request.StatePath}, ErrConfirmationRequired
}
if !versionPattern.MatchString(request.Version) {
return Result{StatePath: request.StatePath}, errors.New("Pi version must be an explicit pinned version")
}
if request.Source == "" {
request.Source = BuildSource
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 != BuildSource && request.Source != PullSource {
return Result{StatePath: request.StatePath}, errors.New("Pi update source must be build or pull")
return Result{StatePath: request.StatePath}, fmt.Errorf("%w: Pi update source must be build or pull", ErrInvalidRequest)
}
if request.Source == PullSource && !digestPattern.MatchString(request.Image) {
return Result{StatePath: request.StatePath}, errors.New("pulled Pi image must use an immutable sha256 digest")
if request.Source == PullSource {
canonical, err := canonicalDigestReference(request.Image)
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 {
return Result{StatePath: request.StatePath}, ErrInterruptedUpdate
} 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 }
defer func() {
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) }
}
}()
running, err := activeSessions(ctx, runner)
if err != nil {
return Result{StatePath: request.StatePath}, err
}
frontendStopped := false
if running {
if !request.Drain {
return Result{StatePath: request.StatePath}, ErrActiveSessions
}
stopped, stopErr := runCompose(ctx, runner, "stop", "frontend")
if stopErr != nil {
return Result{StatePath: request.StatePath}, commandError("frontend admission gate", stopped, stopErr)
}
frontendStopped = true
defer func() {
if frontendStopped {
_, _ = runCompose(context.Background(), runner, "up", "--detach", "frontend")
}
}()
for attempts := 0; attempts < 30; attempts++ {
running, err = activeSessions(ctx, runner)
if err != nil {
@@ -134,41 +133,34 @@ func Update(ctx context.Context, runner Runner, request Request) (Result, error)
}
state.Phase = PhaseBuilding
if err := writeState(request.StatePath, state); err != nil {
return Result{StatePath: request.StatePath}, err
}
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 {
state.Phase, state.Error = PhaseFailed, "candidate image preparation failed"
_ = writeState(request.StatePath, state)
return Result{Phase: PhaseFailed, StatePath: request.StatePath}, err
return compensate(ctx, runner, request.StatePath, state, err)
}
if frontendStopped {
if request.Drain {
running, err = activeSessions(ctx, runner)
if err != nil {
return Result{Phase: PhaseFailed, StatePath: request.StatePath}, err
return compensate(ctx, runner, request.StatePath, state, err)
}
if running {
state.Phase, state.Error = PhaseFailed, "new session admitted while draining"
_ = writeState(request.StatePath, state)
return Result{Phase: PhaseFailed, StatePath: request.StatePath}, ErrActiveSessions
return compensate(ctx, runner, request.StatePath, state, ErrActiveSessions)
}
}
if err := recreateCore(ctx, runner); err != nil {
state.Phase, state.Error = PhaseFailed, "core recreation failed"
_ = writeState(request.StatePath, state)
return Result{Phase: PhaseFailed, StatePath: request.StatePath}, err
return compensate(ctx, runner, request.StatePath, state, err)
}
state.Phase = PhaseRecreated
state.Candidate, _ = runningImage(ctx, runner, previous.Reference)
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 Result{Phase: PhaseRecreated, StatePath: request.StatePath}, err
return compensate(ctx, runner, request.StatePath, state, err)
}
if err := verifyCandidate(ctx, runner, request.Version, previous); err != nil {
return rollbackAfterFailure(ctx, runner, request.StatePath, state, err)
return compensate(ctx, runner, request.StatePath, state, err)
}
state.Phase, state.Error = PhaseVerified, ""
if err := writeState(request.StatePath, state); err != nil {
return Result{Phase: PhaseVerified, StatePath: request.StatePath}, err
return compensate(ctx, runner, request.StatePath, state, err)
}
return Result{Phase: PhaseVerified, StatePath: request.StatePath}, nil
}
@@ -189,25 +181,35 @@ func Rollback(ctx context.Context, runner Runner, statePath string, confirm bool
}
if err := restore(ctx, runner, state.Previous); err != nil {
state.Phase, state.Error = PhaseFailed, "rollback failed"
_ = writeState(statePath, state)
if writeErr := 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 {
return Result{Phase: PhaseRolledBack, StatePath: statePath}, err
return Result{Phase: PhaseFailed, StatePath: statePath}, errors.New("rollback restored the core but recovery state could not be persisted")
}
return Result{Phase: PhaseRolledBack, StatePath: statePath}, nil
}
func rollbackAfterFailure(ctx context.Context, runner Runner, statePath string, state State, cause error) (Result, error) {
func compensate(ctx context.Context, runner Runner, statePath string, state State, cause error) (Result, error) {
if restoreErr := restore(ctx, runner, state.Previous); restoreErr != nil {
state.Phase, state.Error = PhaseFailed, "candidate verification and automatic rollback failed"
_ = writeState(statePath, state)
return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("candidate verification failed; automatic rollback also 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")
}
return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("update failed; automatic rollback also failed: recovery required")
}
state.Phase, state.Error = PhaseRolledBack, ""
_ = writeState(statePath, state)
return Result{Phase: PhaseRolledBack, StatePath: statePath}, fmt.Errorf("candidate verification failed; previous core image was restored")
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")
}
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
}
func sourceValue(request Request) string {
@@ -217,6 +219,29 @@ func sourceValue(request Request) string {
return string(BuildSource)
}
func canonicalDigestReference(value string) (string, error) {
if strings.Contains(value, "://") || strings.ContainsAny(value, "?#") || strings.Contains(value, "@") && strings.Contains(strings.Split(value, "@")[0], ":") && strings.Contains(strings.Split(value, "@")[0], "//") {
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") }
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")
}
return reference.FamiliarString(canonical), nil
}
// The command text is fixed; no operator input or host path is interpolated into the core shell.
// The marker lives alongside SETTINGS_FILE's named/bind-mounted directory and is read by backend.
func setMaintenance(ctx context.Context, runner Runner, enabled bool) error {
command := "mkdir -p /data/settings && : > /data/settings/maintenance.json && chmod 600 /data/settings/maintenance.json"
if !enabled { command = "rm -f /data/settings/maintenance.json" }
result, err := runCompose(ctx, runner, "exec", "-T", "core", "sh", "-ceu", command)
if err != nil { return commandError("maintenance admission gate", result, err) }
return nil
}
func activeSessions(ctx context.Context, runner Runner) (bool, error) {
result, err := runCompose(ctx, runner, "exec", "-T", "core", "tht", "session", "list", "--json")
if err != nil {
@@ -251,20 +276,27 @@ func runningImage(ctx context.Context, runner Runner, reference string) (Image,
if err != nil {
return Image{}, commandError("core volume check", mounts, err)
}
var contract []Mount
if err := json.Unmarshal([]byte(mounts.Stdout), &contract); err != nil {
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"`
}
if err := json.Unmarshal([]byte(mounts.Stdout), &raw); err != nil {
return Image{}, errors.New("core returned invalid persistence mount data")
}
if len(contract) == 0 {
if len(raw) == 0 {
return Image{}, errors.New("core has no persistence mounts to preserve")
}
volumes := make([]string, 0, len(contract))
for _, mount := range contract {
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)
}
}
return Image{ID: strings.TrimSpace(image.Stdout), Reference: reference, Volumes: volumes, Mounts: contract}, nil
return Image{ID: strings.TrimSpace(image.Stdout), Reference: reference, Volumes: volumes, Mounts: contract, MountFingerprint: mountFingerprint(contract)}, nil
}
func prepareCandidate(ctx context.Context, runner Runner, request Request, reference string) error {
@@ -382,7 +414,7 @@ func sameMounts(left, right []Mount) bool {
return false
}
key := func(m Mount) string {
return m.Type + "\x00" + m.Name + "\x00" + m.Source + "\x00" + m.Destination + "\x00" + fmt.Sprint(m.RW)
return m.Type + "\x00" + m.Name + "\x00" + m.SourceSHA256 + "\x00" + m.Destination + "\x00" + fmt.Sprint(m.RW) + "\x00" + m.Options
}
a, b := make([]string, len(left)), make([]string, len(right))
for i := range left {
+22 -8
View File
@@ -3,6 +3,7 @@ package pi
import (
"context"
"errors"
"fmt"
"io"
"os"
"path/filepath"
@@ -103,10 +104,9 @@ func TestUpdateDoesNotRecreateWhenPreflightOrBuildFails(t *testing.T) {
if err == nil {
t.Fatal("Update() error = nil, want failure")
}
if result.Phase == PhaseRolledBack {
t.Fatalf("pre-recreate failure unexpectedly rolled back: %+v", result)
}
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") }
})
}
}
@@ -162,7 +162,7 @@ func TestRollbackRestoresInterruptedOrPreviouslyRecordedState(t *testing.T) {
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"}}})
writeStateForTest(t, statePath, State{Phase: PhaseRecreated, Previous: Image{ID: "sha256:old", Reference: "thothii-core:local", Volumes: []string{"settings"}, MountFingerprint: "recorded"}})
_, 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)
@@ -180,6 +180,17 @@ func TestRunningImageCapturesServerBindAndNamedMountIdentity(t *testing.T) {
if len(image.Mounts) != 3 || image.Mounts[0].Type != "bind" || image.Mounts[0].Destination != "/data" {
t.Fatalf("mounts = %#v", image.Mounts)
}
if strings.Contains(fmt.Sprint(image), "/srv/thothii") || image.Mounts[0].SourceSHA256 == "" || image.MountFingerprint == "" {
t.Fatalf("mount contract leaked a server source or lacks a safe fingerprint: %#v", image)
}
}
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) }
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) }
}
}
type fakeRunner struct {
@@ -188,12 +199,13 @@ type fakeRunner struct {
version string
activeSessions bool
built bool
currentImage string
volumes []string
mountsJSON string
}
func newFakeRunner() *fakeRunner {
return &fakeRunner{version: "0.80.3", volumes: []string{"settings", "pi-state", "sessions", "workspace-registry"}}
return &fakeRunner{version: "0.80.3", currentImage: "sha256:old", volumes: []string{"settings", "pi-state", "sessions", "workspace-registry"}}
}
func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose.Result, error) {
@@ -201,6 +213,7 @@ func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose
f.calls = append(f.calls, call)
if 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")
@@ -214,7 +227,7 @@ func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose
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")
}
if f.fail == "smoke" && strings.Contains(call, "curl -fsS http://127.0.0.1:8787/models") {
if f.fail == "smoke" && f.built && strings.Contains(call, "127.0.0.1:8787/models") {
return compose.Result{ExitCode: 1}, errors.New("smoke token=secret")
}
switch {
@@ -223,7 +236,7 @@ func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose
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: "sha256:old\n"}, nil
return compose.Result{Stdout: f.currentImage + "\n"}, nil
case strings.Contains(call, "inspect --format {{json .Mounts}}"):
if f.mountsJSON != "" {
return compose.Result{Stdout: f.mountsJSON}, nil
@@ -238,6 +251,7 @@ func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose
case strings.Contains(call, "compose build"):
f.built = true
f.version = "0.81.0"
f.currentImage = "sha256:candidate"
return compose.Result{}, nil
case strings.Contains(call, "pi --version"):
return compose.Result{Stdout: f.version + "\n"}, nil