refactor(cli): rename operator command to tht

This commit is contained in:
2026-08-15 21:56:40 +02:00
parent 460caa550c
commit aa8a2e9278
49 changed files with 303 additions and 261 deletions
+56
View File
@@ -0,0 +1,56 @@
// Package compose executes Docker Compose through a fixed executable and argument arrays.
package compose
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"os"
"os/exec"
)
// Result is the captured output and process exit code for one Docker invocation.
type Result struct {
Stdout string
Stderr string
ExitCode int
}
// Runner executes the Docker CLI. It never invokes a shell.
type Runner struct {
binary string
}
// NewRunner returns a runner for binary. An empty binary selects docker from PATH.
func NewRunner(binary string) Runner {
if binary == "" {
binary = "docker"
}
return Runner{binary: binary}
}
// Run invokes Docker with the supplied argument array and optional standard input.
func (r Runner) Run(ctx context.Context, args []string, stdin io.Reader) (Result, error) {
command := exec.CommandContext(ctx, r.binary, args...)
command.Stdin = stdin
var stdout, stderr bytes.Buffer
command.Stdout = &stdout
command.Stderr = &stderr
err := command.Run()
result := Result{Stdout: stdout.String(), Stderr: stderr.String()}
if err == nil {
return result, nil
}
var exitError *exec.ExitError
if errors.As(err, &exitError) {
result.ExitCode = exitError.ExitCode()
return result, err
}
if errors.Is(err, exec.ErrNotFound) || errors.Is(err, os.ErrNotExist) {
result.ExitCode = 127
return result, fmt.Errorf("%w: %w", exec.ErrNotFound, err)
}
return result, err
}
+66
View File
@@ -0,0 +1,66 @@
package compose
import (
"context"
"errors"
"os"
"os/exec"
"path/filepath"
"strings"
"testing"
)
func TestRunnerPassesEachArgumentWithoutShellSplitting(t *testing.T) {
t.Parallel()
runner := NewRunner(writeExecutable(t, "#!/bin/sh\nprintf '<%s>\\n' \"$@\"\ncat\n"))
result, err := runner.Run(context.Background(), []string{"compose", "--project-directory", "/tmp/a project with spaces", "config"}, strings.NewReader("stdin value\n"))
if err != nil {
t.Fatalf("Run() error = %v", err)
}
want := "<compose>\n<--project-directory>\n</tmp/a project with spaces>\n<config>\nstdin value\n"
if result.Stdout != want {
t.Errorf("stdout = %q, want %q", result.Stdout, want)
}
if result.ExitCode != 0 {
t.Errorf("ExitCode = %d, want 0", result.ExitCode)
}
}
func TestRunnerReturnsTheChildExitCode(t *testing.T) {
t.Parallel()
runner := NewRunner(writeExecutable(t, "#!/bin/sh\necho unavailable >&2\nexit 42\n"))
result, err := runner.Run(context.Background(), []string{"compose", "ps"}, nil)
if err == nil {
t.Fatal("Run() error = nil, want child exit error")
}
if result.ExitCode != 42 {
t.Errorf("ExitCode = %d, want 42", result.ExitCode)
}
if result.Stderr != "unavailable\n" {
t.Errorf("stderr = %q, want unavailable output", result.Stderr)
}
}
func TestRunnerReportsMissingDocker(t *testing.T) {
t.Parallel()
runner := NewRunner(filepath.Join(t.TempDir(), "docker-does-not-exist"))
result, err := runner.Run(context.Background(), []string{"compose", "version"}, nil)
if !errors.Is(err, exec.ErrNotFound) {
t.Fatalf("Run() error = %v, want exec.ErrNotFound", err)
}
if result.ExitCode != 127 {
t.Errorf("ExitCode = %d, want 127", result.ExitCode)
}
}
func writeExecutable(t *testing.T, contents string) string {
t.Helper()
path := filepath.Join(t.TempDir(), "fake-docker")
if err := os.WriteFile(path, []byte(contents), 0o700); err != nil {
t.Fatal(err)
}
return path
}
+469
View File
@@ -0,0 +1,469 @@
// Package config loads the non-secret, local installation descriptor used by tht.
package config
import (
"bytes"
"crypto/sha256"
"errors"
"fmt"
"io"
"net/url"
"os"
"path/filepath"
"regexp"
"sort"
"strings"
"sync"
"github.com/aritmolab/thothii/tools/tht/internal/safeio"
"github.com/compose-spec/compose-go/v2/dotenv"
"github.com/sirupsen/logrus"
"gopkg.in/yaml.v3"
)
const installationFileName = "thothii-installation.yaml"
const maxEnvironmentFileBytes = 1 << 20
const maxSecretSources = 32
var dotenvParseMu sync.Mutex
type descriptor struct {
Profile string `yaml:"profile"`
ProjectDirectory string `yaml:"projectDirectory"`
EnvFile string `yaml:"envFile"`
WorkspaceRepository workspaceRepositoryDescriptor `yaml:"workspaceRepository"`
Overrides []string `yaml:"overrides"`
}
type workspaceRepositoryDescriptor struct {
Remote string `yaml:"remote"`
Branch string `yaml:"branch"`
Access string `yaml:"access"`
}
// WorkspaceRepository is the non-secret Git source identity declared by one installation.
type WorkspaceRepository struct {
Remote string
Branch string
Access string
}
// Installation is a validated local Compose installation. It intentionally contains paths, not
// environment values or secret content.
type Installation struct {
Path string
Profile string
ProjectDirectory string
EnvFile string
WorkspaceRepository WorkspaceRepository
Overrides []string
}
// Load reads and validates an installation descriptor at an absolute path.
func Load(path string) (Installation, error) {
if !filepath.IsAbs(path) {
return Installation{}, fmt.Errorf("installation path must be absolute")
}
path = filepath.Clean(path)
if filepath.Base(path) != installationFileName {
return Installation{}, fmt.Errorf("installation file must be named %s", installationFileName)
}
if err := requireRegularFile(path, "installation file"); err != nil {
return Installation{}, err
}
file, err := os.Open(path)
if err != nil {
return Installation{}, fmt.Errorf("open installation file: %w", err)
}
defer file.Close()
var raw descriptor
decoder := yaml.NewDecoder(file)
decoder.KnownFields(true)
if err := decoder.Decode(&raw); err != nil {
return Installation{}, fmt.Errorf("read installation file: %w", err)
}
if err := ensureOnlyOneDocument(decoder); err != nil {
return Installation{}, err
}
if raw.Profile != "local" && raw.Profile != "server" {
return Installation{}, fmt.Errorf("profile must be local or server")
}
if err := requireDirectory(raw.ProjectDirectory, "projectDirectory"); err != nil {
return Installation{}, err
}
if err := requireRegularFile(raw.EnvFile, "envFile"); err != nil {
return Installation{}, err
}
installation := Installation{
Path: path,
Profile: raw.Profile,
ProjectDirectory: filepath.Clean(raw.ProjectDirectory),
EnvFile: filepath.Clean(raw.EnvFile),
WorkspaceRepository: WorkspaceRepository{
Remote: raw.WorkspaceRepository.Remote,
Branch: raw.WorkspaceRepository.Branch,
Access: raw.WorkspaceRepository.Access,
},
Overrides: make([]string, 0, len(raw.Overrides)),
}
for _, override := range raw.Overrides {
if err := requireRegularFile(override, "override"); err != nil {
return Installation{}, err
}
installation.Overrides = append(installation.Overrides, filepath.Clean(override))
}
for _, composeFile := range installation.ComposeFiles()[:2] {
if err := requireRegularFile(composeFile, "Compose file"); err != nil {
return Installation{}, err
}
}
if err := installation.validateWorkspaceRepository(); err != nil {
return Installation{}, err
}
if info, err := os.Lstat(installation.CurrentImageOverridePath()); err == nil {
if !info.Mode().IsRegular() {
return Installation{}, errors.New("installation current-image override must be a regular file")
}
} else if !errors.Is(err, os.ErrNotExist) {
return Installation{}, errors.New("installation current-image override could not be inspected")
}
return installation, nil
}
var safeGitBranch = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._/-]*$`)
var scpSSHRemote = regexp.MustCompile(`^git@[^:/\s]+:[^\s]+$`)
func (i Installation) validateWorkspaceRepository() error {
gitAccess := ""
for _, override := range i.Overrides {
switch filepath.Base(override) {
case "compose.git-ssh.yaml":
if gitAccess != "" {
return errors.New("installation must select exactly one Git transport override")
}
gitAccess = "ssh"
case "compose.git-https.yaml":
if gitAccess != "" {
return errors.New("installation must select exactly one Git transport override")
}
gitAccess = "https"
}
}
declared := i.WorkspaceRepository
if gitAccess == "" {
if declared.Remote != "" || declared.Branch != "" || declared.Access != "" {
return errors.New("workspaceRepository requires one Git transport override")
}
return nil
}
if declared.Remote == "" || declared.Branch == "" || declared.Access == "" {
return errors.New("workspaceRepository is required for a Git installation")
}
if declared.Access != gitAccess {
return errors.New("workspaceRepository access does not match the Git transport override")
}
if !safeGitBranch.MatchString(declared.Branch) || strings.Contains(declared.Branch, "..") ||
strings.Contains(declared.Branch, "@{") || strings.HasPrefix(declared.Branch, "-") ||
strings.HasSuffix(declared.Branch, ".lock") {
return errors.New("workspaceRepository branch is invalid")
}
if err := validateRepositoryRemote(declared.Remote, declared.Access); err != nil {
return err
}
values, err := i.environmentValues()
if err != nil {
return err
}
if values["THT_WORKSPACE_GIT_REMOTE"] != declared.Remote ||
values["THT_WORKSPACE_GIT_BRANCH"] != declared.Branch {
return errors.New("workspaceRepository does not match the installation environment")
}
required := []string{"THT_WORKSPACE_GIT_CREDENTIALS_FILE", "THT_WORKSPACE_GIT_CA_FILE"}
if gitAccess == "ssh" {
required = []string{"THT_WORKSPACE_GIT_SSH_KEY_FILE", "THT_WORKSPACE_GIT_KNOWN_HOSTS_FILE"}
}
for _, name := range required {
if err := requireRegularFile(values[name], "workspaceRepository credential"); err != nil {
return errors.New("workspaceRepository credentials are unavailable")
}
}
return nil
}
func validateRepositoryRemote(remote, access string) error {
if remote == "" || strings.TrimSpace(remote) != remote || strings.ContainsRune(remote, '\x00') {
return errors.New("workspaceRepository remote is invalid")
}
if access == "ssh" && scpSSHRemote.MatchString(remote) {
return nil
}
parsed, err := url.Parse(remote)
if err != nil || parsed.Hostname() == "" || parsed.RawQuery != "" || parsed.Fragment != "" ||
parsed.User != nil && access == "https" || parsed.User != nil && strings.Contains(parsed.User.String(), ":") {
return errors.New("workspaceRepository remote is invalid")
}
if access == "https" && parsed.Scheme != "https" {
return errors.New("workspaceRepository remote does not match HTTPS access")
}
if access == "ssh" && parsed.Scheme != "ssh" {
return errors.New("workspaceRepository remote does not match SSH access")
}
return nil
}
// ComposeFiles returns the base file, selected profile file, and declared optional overrides in
// the exact order Compose applies them.
func (i Installation) ComposeFiles() []string {
files := []string{
filepath.Join(i.ProjectDirectory, "compose.yaml"),
filepath.Join(i.ProjectDirectory, "deploy", "compose."+i.Profile+".yaml"),
}
files = append(files, i.Overrides...)
currentImage := i.CurrentImageOverridePath()
if info, err := os.Lstat(currentImage); err == nil && info.Mode().IsRegular() {
files = append(files, currentImage)
}
return files
}
// ControlDirectory contains state that is private to one installation descriptor, even when
// multiple installations intentionally share one source checkout.
func (i Installation) ControlDirectory() string {
return filepath.Join(i.ProjectDirectory, ".tht", i.ProjectName())
}
func (i Installation) CurrentImageOverridePath() string {
return filepath.Join(i.ControlDirectory(), "current-image.yaml")
}
func (i Installation) UpdateStatePath() string {
return filepath.Join(i.ControlDirectory(), "update-state.json")
}
func (i Installation) RestartStatePath() string {
return filepath.Join(i.ControlDirectory(), "restart-state.json")
}
// ProjectName is stable for one installation and avoids collisions between different checkouts.
func (i Installation) ProjectName() string {
sum := sha256.Sum256([]byte(i.Path))
return fmt.Sprintf("thothii-%x", sum[:6])
}
// ComposeArgs builds Docker Compose arguments without shell quoting or interpolation.
func (i Installation) ComposeArgs(command ...string) []string {
return i.composeArgs(i.ComposeFiles(), command...)
}
// ComposeArgsWithFinalOverride appends one validated, generated override after every durable
// installation selector and before the Compose command.
func (i Installation) ComposeArgsWithFinalOverride(override string, command ...string) ([]string, error) {
if filepath.Clean(override) != override || !filepath.IsAbs(override) {
return nil, errors.New("final Compose override must be an absolute canonical path")
}
if err := requireRegularFile(override, "final Compose override"); err != nil {
return nil, err
}
files := append(i.ComposeFiles(), override)
return i.composeArgs(files, command...), nil
}
func (i Installation) composeArgs(files []string, command ...string) []string {
args := []string{"compose", "--project-name", i.ProjectName(), "--project-directory", i.ProjectDirectory, "--env-file", i.EnvFile}
for _, composeFile := range files {
args = append(args, "-f", composeFile)
}
return append(args, command...)
}
// SecretFiles returns canonical local secret paths declared through *_FILE or *_SOURCE variables.
// Compose's dotenv parser resolves comments, quotes, escapes, and interpolation. Unsupported or
// unresolved source interpolation is rejected before tht invokes Docker.
func (i Installation) SecretFiles() ([]string, error) {
contents, err := safeio.ReadCanonicalRegular(i.EnvFile, maxEnvironmentFileBytes)
if err != nil {
return nil, errors.New("installation secret declarations could not be read")
}
values, err := parseComposeDotenv(contents)
if err != nil {
return nil, errors.New("installation secret declarations could not be read")
}
files := make([]string, 0, len(values))
seen := make(map[string]struct{})
for key, value := range values {
key = strings.ToUpper(key)
if !strings.HasSuffix(key, "_FILE") && !strings.HasSuffix(key, "_SOURCE") {
continue
}
if err := safeio.ValidateCanonicalPath(value); err != nil {
return nil, errors.New("installation secret declarations could not be read")
}
if _, exists := seen[value]; !exists {
files = append(files, value)
seen[value] = struct{}{}
if len(files) > maxSecretSources {
return nil, errors.New("installation secret declarations could not be read")
}
}
}
sort.Strings(files)
return files, nil
}
// EnvironmentValue returns one declared installation value without exposing dotenv parsing to
// callers. It is used only for operator-visible file locations, never for secret content.
func (i Installation) EnvironmentValue(name string) (string, error) {
values, err := i.environmentValues()
if err != nil {
return "", err
}
return values[name], nil
}
func (i Installation) environmentValues() (map[string]string, error) {
contents, err := safeio.ReadCanonicalRegular(i.EnvFile, maxEnvironmentFileBytes)
if err != nil {
return nil, errors.New("installation environment could not be read")
}
values, err := parseComposeDotenv(contents)
if err != nil {
return nil, errors.New("installation environment could not be read")
}
return values, nil
}
// PreservationPaths returns the server bind roots, backup root, and declared secret files whose
// filesystem identities must survive a data-preserving removal.
func (i Installation) PreservationPaths() ([]string, error) {
if i.Profile != "server" {
return nil, errors.New("data-preserving removal requires a server installation")
}
values, err := i.environmentValues()
if err != nil {
return nil, err
}
paths := make([]string, 0)
seen := make(map[string]struct{})
for _, name := range []string{
"THT_DATA_ROOT", "THT_PI_STATE_ROOT", "THT_WORKSPACE_REGISTRY_ROOT", "THT_BACKUP_ROOT",
} {
path := values[name]
if err := requireCanonicalDirectory(path); err != nil {
return nil, fmt.Errorf("%s must identify an existing canonical directory", name)
}
if _, exists := seen[path]; !exists {
paths = append(paths, path)
seen[path] = struct{}{}
}
}
secretFiles, err := i.SecretFiles()
if err != nil {
return nil, err
}
for _, path := range secretFiles {
if _, exists := seen[path]; !exists {
paths = append(paths, path)
seen[path] = struct{}{}
}
}
return paths, nil
}
func requireCanonicalDirectory(path string) error {
if err := safeio.ValidateCanonicalPath(path); err != nil {
return err
}
resolved, err := filepath.EvalSymlinks(path)
if err != nil || resolved != path {
return errors.New("directory path is unavailable or contains a symlink")
}
info, err := os.Stat(path)
if err != nil || !info.IsDir() {
return errors.New("directory path is unavailable")
}
return nil
}
func parseComposeDotenv(contents []byte) (map[string]string, error) {
dotenvParseMu.Lock()
defer dotenvParseMu.Unlock()
logger := logrus.StandardLogger()
previousOutput := logger.Out
previousHooks := logger.ReplaceHooks(make(logrus.LevelHooks))
logger.SetOutput(io.Discard)
warnings := &dotenvWarnings{}
logger.AddHook(warnings)
defer func() {
logger.SetOutput(previousOutput)
logger.ReplaceHooks(previousHooks)
}()
values, err := dotenv.ParseWithLookup(bytes.NewReader(contents), os.LookupEnv)
if err != nil || warnings.seen {
return nil, errors.New("dotenv parsing failed")
}
return values, nil
}
type dotenvWarnings struct {
seen bool
}
func (w *dotenvWarnings) Levels() []logrus.Level {
return logrus.AllLevels
}
func (w *dotenvWarnings) Fire(entry *logrus.Entry) error {
if entry.Level == logrus.WarnLevel {
w.seen = true
}
return nil
}
func ensureOnlyOneDocument(decoder *yaml.Decoder) error {
var extra any
err := decoder.Decode(&extra)
if errors.Is(err, io.EOF) {
return nil
}
if err != nil {
return fmt.Errorf("read installation file: %w", err)
}
return fmt.Errorf("installation file must contain one YAML document")
}
func requireDirectory(path, field string) error {
if !filepath.IsAbs(path) {
return fmt.Errorf("%s must be an absolute path", field)
}
info, err := os.Stat(path)
if err != nil {
return fmt.Errorf("%s is unavailable: %w", field, err)
}
if !info.IsDir() {
return fmt.Errorf("%s must be a directory", field)
}
return nil
}
func requireRegularFile(path, field string) error {
if !filepath.IsAbs(path) {
return fmt.Errorf("%s must be an absolute path", field)
}
info, err := os.Stat(path)
if err != nil {
return fmt.Errorf("%s is unavailable: %w", field, err)
}
if !info.Mode().IsRegular() {
return fmt.Errorf("%s must be a regular file", field)
}
return nil
}
@@ -0,0 +1,309 @@
package config
import (
"os"
"path/filepath"
"strings"
"testing"
)
func TestLoadSelectsLocalComposeFilesForAnInstallationInPathsWithSpaces(t *testing.T) {
t.Parallel()
installationPath, projectDirectory, envFile, override := writeInstallation(t, "local")
installation, err := Load(installationPath)
if err != nil {
t.Fatalf("Load() error = %v", err)
}
if installation.ProjectDirectory != projectDirectory {
t.Errorf("ProjectDirectory = %q, want %q", installation.ProjectDirectory, projectDirectory)
}
if installation.EnvFile != envFile {
t.Errorf("EnvFile = %q, want %q", installation.EnvFile, envFile)
}
if !strings.Contains(installationPath, "installation folder with spaces") {
t.Fatalf("test setup must exercise a path with spaces: %q", installationPath)
}
want := []string{
filepath.Join(projectDirectory, "compose.yaml"),
filepath.Join(projectDirectory, "deploy", "compose.local.yaml"),
override,
}
assertStringsEqual(t, installation.ComposeFiles(), want)
}
func TestLoadSelectsServerComposeFiles(t *testing.T) {
t.Parallel()
installationPath, projectDirectory, _, override := writeInstallation(t, "server")
installation, err := Load(installationPath)
if err != nil {
t.Fatalf("Load() error = %v", err)
}
want := []string{
filepath.Join(projectDirectory, "compose.yaml"),
filepath.Join(projectDirectory, "deploy", "compose.server.yaml"),
override,
}
assertStringsEqual(t, installation.ComposeFiles(), want)
}
func TestLoadRequiresAndReturnsTypedWorkspaceRepositoryForGitInstallations(t *testing.T) {
installationPath, projectDirectory, envFile, _ := writeInstallation(t, "local")
gitOverride := filepath.Join(projectDirectory, "deploy", "compose.git-ssh.yaml")
if err := os.WriteFile(gitOverride, []byte("services: {}\n"), 0o600); err != nil {
t.Fatal(err)
}
secretRoot := filepath.Dir(envFile)
privateKey := filepath.Join(secretRoot, "git-key")
knownHosts := filepath.Join(secretRoot, "known-hosts")
for _, file := range []string{privateKey, knownHosts} {
if err := os.WriteFile(file, []byte("fixture\n"), 0o600); err != nil {
t.Fatal(err)
}
}
remote := "git@gitea.example.org:clinical/workspaces.git"
environment := strings.Join([]string{
"THT_WORKSPACE_GIT_REMOTE=" + remote,
"THT_WORKSPACE_GIT_BRANCH=main",
"THT_WORKSPACE_GIT_SSH_KEY_FILE=" + privateKey,
"THT_WORKSPACE_GIT_KNOWN_HOSTS_FILE=" + knownHosts,
}, "\n") + "\n"
if err := os.WriteFile(envFile, []byte(environment), 0o600); err != nil {
t.Fatal(err)
}
contents := "profile: local\nprojectDirectory: " + projectDirectory +
"\nenvFile: " + envFile +
"\nworkspaceRepository:\n remote: " + remote +
"\n branch: main\n access: ssh\noverrides:\n - " + gitOverride + "\n"
if err := os.WriteFile(installationPath, []byte(contents), 0o600); err != nil {
t.Fatal(err)
}
installation, err := Load(installationPath)
if err != nil {
t.Fatal(err)
}
if installation.WorkspaceRepository.Remote != remote ||
installation.WorkspaceRepository.Branch != "main" ||
installation.WorkspaceRepository.Access != "ssh" {
t.Fatalf("WorkspaceRepository = %#v", installation.WorkspaceRepository)
}
}
func TestLoadRejectsGitOverrideWithoutTypedWorkspaceRepository(t *testing.T) {
installationPath, projectDirectory, envFile, _ := writeInstallation(t, "local")
gitOverride := filepath.Join(projectDirectory, "deploy", "compose.git-https.yaml")
if err := os.WriteFile(gitOverride, []byte("services: {}\n"), 0o600); err != nil {
t.Fatal(err)
}
contents := "profile: local\nprojectDirectory: " + projectDirectory +
"\nenvFile: " + envFile + "\noverrides:\n - " + gitOverride + "\n"
if err := os.WriteFile(installationPath, []byte(contents), 0o600); err != nil {
t.Fatal(err)
}
_, err := Load(installationPath)
if err == nil || !strings.Contains(err.Error(), "workspaceRepository") {
t.Fatalf("Load() error = %v, want workspaceRepository error", err)
}
}
func TestComposeArgsAutomaticallyIncludeTheInstallationCurrentImageOverride(t *testing.T) {
t.Parallel()
installationPath, _, _, _ := writeInstallation(t, "local")
seed, err := Load(installationPath)
if err != nil {
t.Fatal(err)
}
currentImage := seed.CurrentImageOverridePath()
if err := os.MkdirAll(filepath.Dir(currentImage), 0o700); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(currentImage, []byte("services:\n core:\n image: candidate\n"), 0o600); err != nil {
t.Fatal(err)
}
installation, err := Load(installationPath)
if err != nil {
t.Fatal(err)
}
args := installation.ComposeArgs("up", "--detach")
want := []string{"-f", currentImage, "up", "--detach"}
if !containsSequence(args, want) {
t.Fatalf("ComposeArgs() = %#v, want durable override immediately before command", args)
}
}
func TestComposeArgsWithFinalOverridePreservesCurrentImagePrecedence(t *testing.T) {
installationPath, _, _, _ := writeInstallation(t, "server")
seed, err := Load(installationPath)
if err != nil {
t.Fatal(err)
}
if err := os.MkdirAll(seed.ControlDirectory(), 0o700); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(seed.CurrentImageOverridePath(), []byte("services: {}\n"), 0o600); err != nil {
t.Fatal(err)
}
final := filepath.Join(seed.ControlDirectory(), "migration.yaml")
if err := os.WriteFile(final, []byte("services: {}\n"), 0o600); err != nil {
t.Fatal(err)
}
installation, err := Load(installationPath)
if err != nil {
t.Fatal(err)
}
args, err := installation.ComposeArgsWithFinalOverride(final, "--profile", "session-migrate", "config")
if err != nil {
t.Fatal(err)
}
want := []string{"-f", installation.CurrentImageOverridePath(), "-f", final, "--profile", "session-migrate", "config"}
if !containsSequence(args, want) {
t.Fatalf("ComposeArgsWithFinalOverride() = %#v, want %#v", args, want)
}
}
func TestPreservationPathsReturnsCanonicalBindRootsBackupsAndSecretFiles(t *testing.T) {
installationPath, _, envFile, _ := writeInstallation(t, "server")
root := filepath.Dir(envFile)
var wanted []string
var lines []string
for _, item := range []struct{ key, name string }{
{"THT_DATA_ROOT", "data"},
{"THT_PI_STATE_ROOT", "pi-state"},
{"THT_WORKSPACE_REGISTRY_ROOT", "workspace-registry"},
{"THT_BACKUP_ROOT", "backups"},
} {
path := filepath.Join(root, item.name)
if err := os.Mkdir(path, 0o700); err != nil {
t.Fatal(err)
}
wanted = append(wanted, path)
lines = append(lines, item.key+"="+path)
}
secret := filepath.Join(root, "secret")
if err := os.WriteFile(secret, []byte("secret"), 0o600); err != nil {
t.Fatal(err)
}
wanted = append(wanted, secret)
lines = append(lines, "APP_TOKEN_FILE="+secret)
if err := os.WriteFile(envFile, []byte(strings.Join(lines, "\n")+"\n"), 0o600); err != nil {
t.Fatal(err)
}
installation, err := Load(installationPath)
if err != nil {
t.Fatal(err)
}
got, err := installation.PreservationPaths()
if err != nil {
t.Fatal(err)
}
assertStringsEqual(t, got, wanted)
}
func TestInstallationControlPathsAreIsolatedForDescriptorsSharingOneCheckout(t *testing.T) {
projectDirectory := t.TempDir()
first := Installation{Path: filepath.Join(t.TempDir(), installationFileName), ProjectDirectory: projectDirectory}
second := Installation{Path: filepath.Join(t.TempDir(), installationFileName), ProjectDirectory: projectDirectory}
if first.CurrentImageOverridePath() == second.CurrentImageOverridePath() {
t.Fatalf("shared-checkout installations reused %q", first.CurrentImageOverridePath())
}
for _, installation := range []Installation{first, second} {
if filepath.Dir(filepath.Dir(installation.CurrentImageOverridePath())) != filepath.Join(projectDirectory, ".tht") {
t.Fatalf("current-image path %q is not installation-specific under .tht", installation.CurrentImageOverridePath())
}
if filepath.Dir(installation.UpdateStatePath()) != filepath.Dir(installation.CurrentImageOverridePath()) {
t.Fatalf("state %q and selector %q do not share one installation control directory", installation.UpdateStatePath(), installation.CurrentImageOverridePath())
}
if got, want := installation.RestartStatePath(), filepath.Join(installation.ControlDirectory(), "restart-state.json"); got != want {
t.Fatalf("RestartStatePath() = %q, want %q", got, want)
}
}
}
func TestLoadRejectsRelativeInstallationPaths(t *testing.T) {
t.Parallel()
_, err := Load("thothii-installation.yaml")
if err == nil || !strings.Contains(err.Error(), "absolute") {
t.Fatalf("Load() error = %v, want an absolute-path error", err)
}
}
func writeInstallation(t *testing.T, profile string) (string, string, string, string) {
t.Helper()
temporaryRoot, err := filepath.EvalSymlinks(os.TempDir())
if err != nil {
t.Fatal(err)
}
physicalRoot, err := os.MkdirTemp(temporaryRoot, "tht-config-")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.RemoveAll(physicalRoot) })
root := filepath.Join(physicalRoot, "installation folder with spaces")
projectDirectory := filepath.Join(root, "project directory with spaces")
if err := os.MkdirAll(filepath.Join(projectDirectory, "deploy"), 0o755); err != nil {
t.Fatal(err)
}
for _, name := range []string{"compose.yaml", filepath.Join("deploy", "compose.local.yaml"), filepath.Join("deploy", "compose.server.yaml")} {
if err := os.WriteFile(filepath.Join(projectDirectory, name), []byte("services: {}\n"), 0o600); err != nil {
t.Fatal(err)
}
}
envFile := filepath.Join(root, "environment file.env")
if err := os.WriteFile(envFile, []byte("SAFE_VALUE=1\n"), 0o600); err != nil {
t.Fatal(err)
}
override := filepath.Join(root, "extra override.yaml")
if err := os.WriteFile(override, []byte("services: {}\n"), 0o600); err != nil {
t.Fatal(err)
}
installationPath := filepath.Join(root, "thothii-installation.yaml")
contents := "profile: " + profile + "\nprojectDirectory: " + projectDirectory + "\nenvFile: " + envFile + "\noverrides:\n - " + override + "\n"
if err := os.WriteFile(installationPath, []byte(contents), 0o600); err != nil {
t.Fatal(err)
}
return installationPath, projectDirectory, envFile, override
}
func assertStringsEqual(t *testing.T, got, want []string) {
t.Helper()
if len(got) != len(want) {
t.Fatalf("length = %d, want %d: got %#v", len(got), len(want), got)
}
for i := range want {
if got[i] != want[i] {
t.Errorf("value[%d] = %q, want %q", i, got[i], want[i])
}
}
}
func containsSequence(values, wanted []string) bool {
for start := range values {
if len(values)-start < len(wanted) {
continue
}
matched := true
for offset := range wanted {
if values[start+offset] != wanted[offset] {
matched = false
break
}
}
if matched {
return true
}
}
return false
}
+198
View File
@@ -0,0 +1,198 @@
// Package output removes credentials from diagnostics before they reach an operator terminal.
package output
import (
"bytes"
"encoding/json"
"errors"
"io"
"regexp"
"sort"
"strings"
"github.com/aritmolab/thothii/tools/tht/internal/safeio"
"github.com/compose-spec/compose-go/v2/dotenv"
)
var credentialField = regexp.MustCompile(`(?im)((?:"|')?[\w.-]*(?:password|token|key|secret|credential)[\w.-]*(?:"|')?\s*[:=]\s*)(?:"(?:\\.|[^"\\\r\n])*"|'[^'\r\n]*'|[^\s,;}]+)`)
var dotenvAssignment = regexp.MustCompile(`(?m)^\s*(?:export\s+)?[A-Za-z_][A-Za-z0-9_.-]*\s*=`)
const maxSecretFileBytes = 64 * 1024
const maxSecretSourceFiles = 32
const maxSecretSourceBytes = 256 * 1024
const maxSecretValuesPerFile = 1024
const maxSecretValues = 4096
const maxJSONSecretDepth = 32
const maxDiagnosticDetailBytes = 512
// Sanitize redacts common credential fields and every supplied secret value.
func Sanitize(text string, secretValues []string) string {
text = credentialField.ReplaceAllString(text, "${1}[REDACTED]")
values := append([]string(nil), secretValues...)
sort.Slice(values, func(i, j int) bool { return len(values[i]) > len(values[j]) })
for _, value := range values {
if value != "" {
text = strings.ReplaceAll(text, value, "[REDACTED]")
if encoded, err := json.Marshal(value); err == nil {
text = strings.ReplaceAll(text, string(encoded), "[REDACTED]")
}
}
}
return text
}
// SanitizeDetail redacts the complete subprocess detail before normalizing and bounding the text
// that may be displayed at the CLI boundary.
func SanitizeDetail(text string, secretValues []string) string {
detail := strings.Join(strings.Fields(Sanitize(text, secretValues)), " ")
if len(detail) <= maxDiagnosticDetailBytes {
return detail
}
var bounded strings.Builder
for _, character := range detail {
encoded := string(character)
if bounded.Len()+len(encoded) > maxDiagnosticDetailBytes {
break
}
bounded.WriteString(encoded)
}
return bounded.String()
}
// SecretValuesFromFiles reads non-empty secret-file contents without exposing them to callers.
func SecretValuesFromFiles(paths []string) ([]string, error) {
if len(paths) > maxSecretSourceFiles {
return nil, errors.New("declared secret file could not be read")
}
values := make([]string, 0, len(paths))
seen := make(map[string]struct{})
var totalBytes int64
for _, path := range paths {
contents, size, err := readSecretFile(path)
if err != nil {
return nil, err
}
totalBytes += size
if totalBytes > maxSecretSourceBytes {
return nil, errors.New("declared secret file could not be read")
}
extracted, err := extractSecretValues(contents)
if err != nil {
return nil, errors.New("declared secret file could not be read")
}
for _, value := range extracted {
if value == "" {
continue
}
if _, exists := seen[value]; exists {
continue
}
values = append(values, value)
seen[value] = struct{}{}
if len(values) > maxSecretValues {
return nil, errors.New("declared secret file could not be read")
}
}
}
return values, nil
}
func readSecretFile(path string) ([]byte, int64, error) {
contents, err := safeio.ReadCanonicalRegular(path, maxSecretFileBytes)
if err != nil {
return nil, 0, errors.New("declared secret file could not be read")
}
return contents, int64(len(contents)), nil
}
func extractSecretValues(contents []byte) ([]string, error) {
whole := strings.TrimRight(string(contents), "\r\n")
trimmed := bytes.TrimSpace(contents)
if len(trimmed) == 0 {
return nil, nil
}
values := make([]string, 0, 8)
if whole != "" {
values = append(values, whole)
}
if trimmed[0] == '{' || trimmed[0] == '[' {
var document any
decoder := json.NewDecoder(bytes.NewReader(trimmed))
decoder.UseNumber()
if err := decoder.Decode(&document); err != nil {
return nil, err
}
var extra any
if err := decoder.Decode(&extra); !errors.Is(err, io.EOF) {
if err == nil {
return nil, errors.New("secret JSON contains multiple documents")
}
return nil, err
}
count := 0
if err := collectJSONSecretValues(document, 0, &count, &values); err != nil {
return nil, err
}
return values, nil
}
if dotenvAssignment.Match(trimmed) {
parsed, err := dotenv.Parse(bytes.NewReader(contents))
if err != nil || len(parsed) > maxSecretValuesPerFile {
return nil, errors.New("secret dotenv bundle is invalid")
}
keys := make([]string, 0, len(parsed))
for key := range parsed {
keys = append(keys, key)
}
sort.Strings(keys)
for _, key := range keys {
if parsed[key] != "" {
values = append(values, parsed[key])
}
}
}
return values, nil
}
func collectJSONSecretValues(value any, depth int, count *int, values *[]string) error {
if depth > maxJSONSecretDepth {
return errors.New("secret JSON nesting is too deep")
}
switch typed := value.(type) {
case map[string]any:
keys := make([]string, 0, len(typed))
for key := range typed {
keys = append(keys, key)
}
sort.Strings(keys)
for _, key := range keys {
if err := collectJSONSecretValues(typed[key], depth+1, count, values); err != nil {
return err
}
}
case []any:
for _, item := range typed {
if err := collectJSONSecretValues(item, depth+1, count, values); err != nil {
return err
}
}
default:
*count++
if *count > maxSecretValuesPerFile {
return errors.New("secret JSON contains too many scalar values")
}
if scalar, ok := typed.(string); ok && scalar != "" {
*values = append(*values, scalar)
}
}
return nil
}
+170
View File
@@ -0,0 +1,170 @@
package output
import (
"os"
"path/filepath"
"strings"
"testing"
)
func TestSanitizeRedactsPasswordTokenAndKeyFields(t *testing.T) {
t.Parallel()
got := Sanitize("DB_PASSWORD=hunter2\naccess_token: abc123\napi-key = quoted-value\nplain=safe\n", nil)
want := "DB_PASSWORD=[REDACTED]\naccess_token: [REDACTED]\napi-key = [REDACTED]\nplain=safe\n"
if got != want {
t.Errorf("Sanitize() = %q, want %q", got, want)
}
}
func TestSanitizeRedactsSecretFileContents(t *testing.T) {
t.Parallel()
secretFile := filepath.Join(physicalTempDir(t), "provider-token")
if err := os.WriteFile(secretFile, []byte("top-secret-value\n"), 0o600); err != nil {
t.Fatal(err)
}
secrets, err := SecretValuesFromFiles([]string{secretFile})
if err != nil {
t.Fatalf("SecretValuesFromFiles() error = %v", err)
}
got := Sanitize("request failed for top-secret-value", secrets)
if got != "request failed for [REDACTED]" {
t.Errorf("Sanitize() = %q, want redacted secret", got)
}
}
func TestSecretValuesFromFilesRedactsNestedJSONScalars(t *testing.T) {
t.Parallel()
secretFile := filepath.Join(physicalTempDir(t), "pi-auth.json")
contents := `{
"providers": {
"dummy-provider": {
"auth": {
"key": "dummy-canary-json-key",
"tokens": {
"access": "dummy-canary-json-access",
"refresh": "dummy-canary-json-refresh"
}
}
}
}
}`
if err := os.WriteFile(secretFile, []byte(contents), 0o600); err != nil {
t.Fatal(err)
}
secrets, err := SecretValuesFromFiles([]string{secretFile})
if err != nil {
t.Fatalf("SecretValuesFromFiles() error = %v", err)
}
got := Sanitize(
"unlabelled dummy-canary-json-key dummy-canary-json-access dummy-canary-json-refresh",
secrets,
)
for _, canary := range []string{
"dummy-canary-json-key",
"dummy-canary-json-access",
"dummy-canary-json-refresh",
} {
if strings.Contains(got, canary) {
t.Fatalf("Sanitize() exposed nested JSON scalar %q: %q", canary, got)
}
}
}
func TestSecretValuesFromFilesRedactsEveryDotenvBundleValue(t *testing.T) {
t.Parallel()
secretFile := filepath.Join(physicalTempDir(t), "thothii.secrets")
contents := "MODEL_API_KEY=dummy-canary-bundle-model\n" +
"DWH_PASSWORD='dummy-canary-bundle-dwh'\n" +
"SESSION_TOKEN=\"dummy-canary-bundle-session\"\n"
if err := os.WriteFile(secretFile, []byte(contents), 0o600); err != nil {
t.Fatal(err)
}
secrets, err := SecretValuesFromFiles([]string{secretFile})
if err != nil {
t.Fatalf("SecretValuesFromFiles() error = %v", err)
}
got := Sanitize(
"unlabelled dummy-canary-bundle-model dummy-canary-bundle-dwh dummy-canary-bundle-session",
secrets,
)
for _, canary := range []string{
"dummy-canary-bundle-model",
"dummy-canary-bundle-dwh",
"dummy-canary-bundle-session",
} {
if strings.Contains(got, canary) {
t.Fatalf("Sanitize() exposed dotenv bundle scalar %q: %q", canary, got)
}
}
}
func TestSanitizeRecognizesQuotedCredentialKeys(t *testing.T) {
t.Parallel()
got := Sanitize(`{"key":"dummy-canary-quoted-key","safe":"visible"}`, nil)
if strings.Contains(got, "dummy-canary-quoted-key") || !strings.Contains(got, `"safe":"visible"`) {
t.Fatalf("Sanitize() = %q, want only the quoted credential field redacted", got)
}
}
func TestSecretValuesFromFilesRejectsMalformedJSONAndExcessiveScalars(t *testing.T) {
t.Parallel()
t.Run("malformed", func(t *testing.T) {
secretFile := filepath.Join(physicalTempDir(t), "malformed-auth.json")
if err := os.WriteFile(secretFile, []byte(`{"auth":{"key":"dummy-canary-malformed"}`), 0o600); err != nil {
t.Fatal(err)
}
if _, err := SecretValuesFromFiles([]string{secretFile}); err == nil {
t.Fatal("SecretValuesFromFiles() error = nil, want malformed-JSON failure")
}
})
t.Run("scalar bound", func(t *testing.T) {
secretFile := filepath.Join(physicalTempDir(t), "many-auth-values.json")
values := make([]string, 1025)
for index := range values {
values[index] = `"dummy-canary-value"`
}
if err := os.WriteFile(secretFile, []byte("["+strings.Join(values, ",")+"]"), 0o600); err != nil {
t.Fatal(err)
}
if _, err := SecretValuesFromFiles([]string{secretFile}); err == nil {
t.Fatal("SecretValuesFromFiles() error = nil, want scalar-count failure")
}
})
}
func TestSecretValuesFromFilesRejectsOversizedFiles(t *testing.T) {
t.Parallel()
secretFile := filepath.Join(physicalTempDir(t), "oversized-token")
if err := os.WriteFile(secretFile, make([]byte, 64*1024+1), 0o600); err != nil {
t.Fatal(err)
}
if _, err := SecretValuesFromFiles([]string{secretFile}); err == nil {
t.Fatal("SecretValuesFromFiles() error = nil, want oversized-file error")
}
}
func physicalTempDir(t *testing.T) string {
t.Helper()
root, err := filepath.EvalSymlinks(os.TempDir())
if err != nil {
t.Fatal(err)
}
directory, err := os.MkdirTemp(root, "tht-output-test-")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.RemoveAll(directory) })
return directory
}
+365
View File
@@ -0,0 +1,365 @@
package pi
import (
"bytes"
"context"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"regexp"
"strings"
"github.com/aritmolab/thothii/tools/tht/internal/compose"
)
var choicePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._/-]{0,127}$`)
type Defaults struct {
Provider string `json:"provider"`
Model string `json:"model"`
Thinking string `json:"thinking"`
}
type ModelOption struct {
Provider string `json:"provider"`
ID string `json:"id"`
}
type piOptions struct {
Providers []string `json:"providers"`
Models []ModelOption `json:"models"`
Reasoning []string `json:"reasoning"`
}
type settingsFileSnapshot struct {
Exists bool `json:"exists"`
RawBase64 string `json:"rawBase64"`
}
var internalIdentityHeaders = []string{
"-H", "x-thoth-principal-issuer: tht",
"-H", "x-thoth-principal-subject: tht-maintenance",
"-H", "x-thoth-principal-display-name: Tht 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")
}
before, err := renderedCore(ctx, runner)
if err != nil {
return err
}
options, err := configurationOptions(ctx, runner)
if err != nil {
return err
}
found := false
for _, model := range options.Models {
if model.Provider == value.Provider && model.ID == value.Model {
found = true
}
}
if !found {
return errors.New("provider/model is not in Pi options")
}
thinkingFound := false
for _, reasoning := range options.Reasoning {
if reasoning == value.Thinking {
thinkingFound = true
}
}
if !thinkingFound {
return errors.New("thinking is not in Pi options")
}
old, err := captureSettingsFile(ctx, runner)
if err != nil {
return err
}
oldEffective, err := readEffectiveSettings(ctx, runner)
if err != nil {
return err
}
restore := func(cause error) error {
if restoreErr := restoreSettingsFile(context.Background(), runner, old); restoreErr != nil {
return fmt.Errorf("%w; previous Pi settings restoration could not be verified: %w", cause, restoreErr)
}
restoredEffective, restoreErr := readEffectiveSettings(context.Background(), runner)
if restoreErr != nil || !bytes.Equal(restoredEffective, oldEffective) {
return fmt.Errorf("%w; previous effective Pi settings 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 := readEffectiveSettings(ctx, runner)
if err != nil {
return restore(err)
}
var saved Defaults
if json.Unmarshal(settings, &saved) != nil || saved.Provider != value.Provider || saved.Model != value.Model || saved.Thinking != value.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 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) {
options, err := configurationOptions(ctx, runner)
if err != nil {
return nil, err
}
return options.Models, nil
}
func configurationOptions(ctx context.Context, runner Runner) (piOptions, error) {
args := append([]string{"exec", "-T", "core", "curl", "-fsS"}, internalIdentityHeaders...)
args = append(args, "http://127.0.0.1:8787/pi-management/options")
result, err := runCompose(ctx, runner, args...)
if err != nil {
return piOptions{}, commandError("Pi options check", result, err)
}
var payload piOptions
if json.Unmarshal([]byte(result.Stdout), &payload) != nil || len(payload.Providers) == 0 || len(payload.Models) == 0 || len(payload.Reasoning) == 0 {
return piOptions{}, errors.New("Pi options response is invalid or empty")
}
providers := make(map[string]bool, len(payload.Providers))
for _, provider := range payload.Providers {
if !choicePattern.MatchString(provider) || providers[provider] {
return piOptions{}, errors.New("Pi options response contains an invalid provider")
}
providers[provider] = true
}
models := make(map[string]bool, len(payload.Models))
for _, option := range payload.Models {
key := option.Provider + "\x00" + option.ID
if !providers[option.Provider] || !choicePattern.MatchString(option.ID) || models[key] {
return piOptions{}, errors.New("Pi options response contains an invalid provider/model")
}
models[key] = true
}
reasoning := make(map[string]bool, len(payload.Reasoning))
for _, value := range payload.Reasoning {
if (value != "low" && value != "medium" && value != "high") || reasoning[value] {
return piOptions{}, errors.New("Pi options response contains an invalid reasoning choice")
}
reasoning[value] = true
}
return payload, 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 captureSettingsFile(ctx context.Context, runner Runner) (settingsFileSnapshot, error) {
result, err := runCompose(ctx, runner, "exec", "-T", "core", "node", "/app/backend/dist/settings/settings-cli.js", "--snapshot")
if err != nil {
return settingsFileSnapshot{}, commandError("Pi installation settings snapshot", result, err)
}
var snapshot settingsFileSnapshot
if json.Unmarshal([]byte(result.Stdout), &snapshot) != nil {
return settingsFileSnapshot{}, errors.New("Pi installation settings snapshot is invalid")
}
raw, decodeErr := base64.StdEncoding.DecodeString(snapshot.RawBase64)
if decodeErr != nil || base64.StdEncoding.EncodeToString(raw) != snapshot.RawBase64 || (!snapshot.Exists && len(raw) != 0) {
return settingsFileSnapshot{}, errors.New("Pi installation settings snapshot is invalid")
}
return snapshot, nil
}
func restoreSettingsFile(ctx context.Context, runner Runner, snapshot settingsFileSnapshot) error {
payload, err := json.Marshal(snapshot)
if err != nil {
return errors.New("Pi installation settings snapshot could not be encoded")
}
args := []string{"compose", "exec", "-T", "core", "node", "/app/backend/dist/settings/settings-cli.js", "--restore"}
result, restoreErr := runner.Run(ctx, args, bytes.NewReader(payload))
verified, verifyErr := captureSettingsFile(ctx, runner)
if restoreErr != nil {
cause := commandError("Pi installation settings restore", result, restoreErr)
if verifyErr == nil && verified == snapshot {
return recoveryRequired("previous Pi settings bytes were restored but durability was not acknowledged", cause)
}
return cause
}
if verifyErr == nil && verified == snapshot {
return nil
}
return errors.New("Pi installation settings restore did not reproduce the exact prior file state")
}
func readEffectiveSettings(ctx context.Context, runner Runner) ([]byte, error) {
args := append([]string{"exec", "-T", "core", "curl", "-fsS"}, internalIdentityHeaders...)
args = append(args, "http://127.0.0.1:8787/settings")
result, err := runCompose(ctx, runner, args...)
if err != nil {
return nil, commandError("Pi installation settings read-back", result, err)
}
var settings map[string]json.RawMessage
if json.Unmarshal([]byte(result.Stdout), &settings) != nil || settings == nil {
return nil, errors.New("Pi installation settings read-back is invalid")
}
canonical, err := json.Marshal(settings)
if err != nil {
return nil, errors.New("Pi installation settings read-back could not be normalized")
}
return canonical, nil
}
// Runner is the narrow, shell-free command boundary shared with tht.
type Runner interface {
Run(context.Context, []string, io.Reader) (compose.Result, error)
}
// Status reports the image-bundled Pi version without using a host Pi executable.
func Status(ctx context.Context, runner Runner) (string, error) {
result, err := runCompose(ctx, runner, "exec", "-T", "core", "pi", "--version")
if err != nil {
return "", commandError("Pi version check", result, err)
}
version := strings.TrimSpace(result.Stdout)
if version == "" {
return "", errors.New("Pi version check returned no version")
}
return version, nil
}
// Doctor verifies the installation-side invariants Pi needs before an update.
func Doctor(ctx context.Context, runner Runner) error {
if _, err := renderedCore(ctx, runner); err != nil {
return err
}
actual, err := Status(ctx, runner)
if err != nil {
return err
}
expected, label, err := expectedVersions(ctx, runner)
if err != nil {
return err
}
if actual != expected || actual != label {
return errors.New("Pi version does not match the image PI_VERSION and io.thothii.pi.version contract")
}
for _, check := range [][]string{
{"exec", "-T", "core", "sh", "-ceu", "test -w /home/thoth/.pi"},
{"exec", "-T", "core", "sh", "-ceu", "test -r /home/thoth/.pi/agent/auth.json"},
{"exec", "-T", "core", "curl", "-fsS", "http://127.0.0.1:8787/health"},
} {
result, err := runCompose(ctx, runner, check...)
if err != nil {
return commandError("Pi preflight check", result, err)
}
}
return Test(ctx, runner)
}
func expectedVersions(ctx context.Context, runner Runner) (string, string, error) {
environment, err := runCompose(ctx, runner, "exec", "-T", "core", "sh", "-ceu", `printf '%s\n' "${PI_VERSION:-}"`)
if err != nil {
return "", "", commandError("Pi expected-version check", environment, err)
}
container, err := runCompose(ctx, runner, "ps", "-q", "core")
if err != nil || strings.TrimSpace(container.Stdout) == "" {
return "", "", commandError("Pi image-label check", container, err)
}
label, err := runner.Run(ctx, []string{"inspect", "--format", `{{ index .Config.Labels "io.thothii.pi.version" }}`, strings.TrimSpace(container.Stdout)}, nil)
if err != nil {
return "", "", commandError("Pi image-label check", label, err)
}
expectedValue, labelValue := strings.TrimSpace(environment.Stdout), strings.TrimSpace(label.Stdout)
if expectedValue == "" || labelValue == "" {
return "", "", errors.New("Pi image expected-version contract is empty")
}
return expectedValue, labelValue, nil
}
// Test retains the direct image-version signal, then delegates all Pi configuration/provider smoke
// validation to core's dedicated, admin-only Pi Management endpoint.
func Test(ctx context.Context, runner Runner) error {
if _, err := Status(ctx, runner); err != nil {
return err
}
args := append([]string{"exec", "-T", "core", "curl", "-fsS", "-X", "POST"}, internalIdentityHeaders...)
args = append(args, "http://127.0.0.1:8787/pi-management/test")
smoke, err := runCompose(ctx, runner, args...)
if err != nil {
return commandError("Pi smoke check", smoke, err)
}
var smokePayload struct {
Ready bool `json:"ready"`
}
if json.Unmarshal([]byte(smoke.Stdout), &smokePayload) != nil || !smokePayload.Ready {
return errors.New("Pi smoke response is not ready")
}
return nil
}
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 map[string]any
if err := json.Unmarshal([]byte(result.Stdout), &document); err != nil {
return Image{}, errors.New("Compose returned invalid rendered configuration")
}
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")
}
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")
}
// 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) {
return runner.Run(ctx, append([]string{"compose"}, args...), nil)
}
func commandError(label string, result compose.Result, err error) error {
if result.ExitCode != 0 {
return commandFailure{message: fmt.Sprintf("%s failed (exit %d)", label, result.ExitCode), exitCode: result.ExitCode}
}
return commandFailure{message: fmt.Sprintf("%s failed", label)}
}
type commandFailure struct {
message string
exitCode int
}
func (e commandFailure) Error() string { return e.message }
// ExitCode exposes a Docker child exit code without exposing its output.
func (e commandFailure) ExitCode() int { return e.exitCode }
+282
View File
@@ -0,0 +1,282 @@
package pi
import (
"context"
"encoding/base64"
"encoding/json"
"errors"
"io"
"strings"
"testing"
"github.com/aritmolab/thothii/tools/tht/internal/compose"
)
func TestDoctorRequiresExternalEndpointAuthPiStateAndHealth(t *testing.T) {
fake := newFakeRunner()
if err := Doctor(context.Background(), fake); err != nil {
t.Fatalf("Doctor() error = %v", err)
}
for _, command := range []string{"pi --version", "PI_VERSION", "io.thothii.pi.version", "test -w /home/thoth/.pi", "test -r /home/thoth/.pi/agent/auth.json", "/health"} {
assertCalled(t, fake.calls, command)
}
}
func TestDoctorRejectsActualEnvironmentAndImageLabelVersionMismatches(t *testing.T) {
for _, mismatch := range []string{"actual", "environment", "label"} {
t.Run(mismatch, func(t *testing.T) {
fake := newFakeRunner()
switch mismatch {
case "actual":
fake.version = "0.80.2"
case "environment":
fake.expectedVersion = "0.80.2"
case "label":
fake.labelVersion = "0.80.2"
}
if err := Doctor(context.Background(), fake); err == nil || !strings.Contains(err.Error(), "version") {
t.Fatalf("Doctor() error = %v, want expected-version mismatch", err)
}
})
}
}
func TestConfigureRestoresAndVerifiesOldSettingsAfterEveryPostSnapshotFailure(t *testing.T) {
for _, failure := range []string{"helper", "readback", "digest"} {
t.Run(failure, func(t *testing.T) {
old := Defaults{Provider: "old", Model: "old-model", Thinking: "low"}
raw, _ := json.Marshal(old)
fake := &configureRunner{failure: failure, settings: old, settingsExist: true, settingsRaw: raw}
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)
}
if !fake.settingsExist || string(fake.settingsRaw) != string(raw) {
t.Fatalf("settings raw snapshot after failure = exists:%t raw:%q, want %q", fake.settingsExist, fake.settingsRaw, raw)
}
})
}
}
func TestSettingsRestoreDoesNotMaskExplicitDurabilityFailureWithMatchingReadback(t *testing.T) {
old := Defaults{Provider: "old", Model: "old-model", Thinking: "low"}
raw, _ := json.Marshal(old)
fake := &configureRunner{
failure: "restore-durability",
settings: Defaults{Provider: "new", Model: "new-model", Thinking: "high"},
settingsExist: true,
settingsRaw: []byte(`{"provider":"new","model":"new-model","thinking":"high"}`),
}
snapshot := settingsFileSnapshot{Exists: true, RawBase64: base64.StdEncoding.EncodeToString(raw)}
err := restoreSettingsFile(context.Background(), fake, snapshot)
var recovery interface{ RecoveryRequired() bool }
if err == nil || !errors.As(err, &recovery) || !recovery.RecoveryRequired() {
t.Fatalf("restore error = %v; want typed recovery-required result", err)
}
if !fake.settingsExist || string(fake.settingsRaw) != string(raw) || fake.settings != old {
t.Fatalf("restored state = exists:%t raw:%q value:%#v; want exact old bytes", fake.settingsExist, fake.settingsRaw, fake.settings)
}
}
func TestConfigurePreservesTypedRecoveryRequiredErrorFromSettingsRestore(t *testing.T) {
old := Defaults{Provider: "old", Model: "old-model", Thinking: "low"}
raw, _ := json.Marshal(old)
fake := &configureRunner{
failure: "helper",
restoreDurabilityFailure: true,
settings: old,
settingsExist: true,
settingsRaw: raw,
}
err := Configure(context.Background(), fake, Defaults{Provider: "new", Model: "new-model", Thinking: "high"})
var recovery interface{ RecoveryRequired() bool }
if err == nil || !errors.As(err, &recovery) || !recovery.RecoveryRequired() {
t.Fatalf("Configure() error = %v; want typed recovery-required result", err)
}
}
type configureRunner struct {
failure string
restoreDurabilityFailure bool
settings Defaults
settingsExist bool
settingsRaw []byte
settingsReads int
configReads int
writes int
}
func (f *configureRunner) Run(_ context.Context, args []string, stdin 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, "/pi-management/options"):
return compose.Result{Stdout: `{"providers":["old","new"],"models":[{"provider":"old","id":"old-model"},{"provider":"new","id":"new-model"}],"reasoning":["low","medium","high"]}`}, nil
case strings.Contains(call, "settings-cli.js --snapshot"):
raw := f.settingsRaw
payload := map[string]any{"exists": f.settingsExist, "rawBase64": base64.StdEncoding.EncodeToString(raw)}
contents, _ := json.Marshal(payload)
return compose.Result{Stdout: string(contents)}, nil
case strings.Contains(call, "settings-cli.js --restore"):
var payload struct {
Exists bool `json:"exists"`
RawBase64 string `json:"rawBase64"`
}
contents, _ := io.ReadAll(stdin)
if json.Unmarshal(contents, &payload) != nil {
return compose.Result{ExitCode: 2}, errors.New("invalid restore payload")
}
f.settingsExist = payload.Exists
f.settingsRaw, _ = base64.StdEncoding.DecodeString(payload.RawBase64)
f.settings = Defaults{Provider: "old", Model: "old-model", Thinking: "low"}
if payload.Exists {
_ = json.Unmarshal(f.settingsRaw, &f.settings)
}
if f.failure == "restore-durability" || f.restoreDurabilityFailure {
return compose.Result{ExitCode: 2}, errors.New("injected post-rename directory fsync failure")
}
return compose.Result{}, nil
case strings.Contains(call, "settings-cli.js"):
if strings.Contains(call, "--provider new") {
f.settings = Defaults{Provider: "new", Model: "new-model", Thinking: "high"}
f.settingsExist = true
f.settingsRaw, _ = json.MarshalIndent(f.settings, "", " ")
f.writes++
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"}
f.settingsExist = true
f.settingsRaw, _ = json.Marshal(f.settings)
}
return compose.Result{}, nil
case strings.Contains(call, "/settings"):
f.settingsReads++
if f.failure == "readback" && f.settings.Provider == "new" {
return compose.Result{Stdout: `{}`}, nil
}
if !f.settingsExist {
return compose.Result{Stdout: `{"provider":"old","model":"old-model","thinking":"low"}`}, nil
}
contents, _ := json.Marshal(f.settings)
return compose.Result{Stdout: string(contents)}, nil
default:
return compose.Result{}, nil
}
}
func TestConfigureAllowsAFirstRunWithoutAnExistingSettingsFile(t *testing.T) {
fake := &configureRunner{}
if err := Configure(context.Background(), fake, Defaults{Provider: "new", Model: "new-model", Thinking: "high"}); err != nil {
t.Fatalf("Configure() clean install error = %v", err)
}
if !fake.settingsExist || fake.writes != 1 || fake.settings.Provider != "new" {
t.Fatalf("clean settings = exists:%t writes:%d value:%#v", fake.settingsExist, fake.writes, fake.settings)
}
}
func TestConfigureCompensationRestoresAbsentAndExactEmptyPriorFiles(t *testing.T) {
for _, prior := range []struct {
name string
exists bool
raw []byte
}{
{name: "absent"},
{name: "empty", exists: true, raw: []byte{}},
{name: "exact raw", exists: true, raw: []byte("{\n \"workspace\": \"kept\",\n \"provider\": \"old\",\n \"model\": \"old-model\",\n \"thinking\": \"low\"\n}\n")},
} {
t.Run(prior.name, func(t *testing.T) {
fake := &configureRunner{failure: "digest", settingsExist: prior.exists, settingsRaw: append([]byte{}, prior.raw...), 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 compensated digest failure")
}
if fake.writes != 1 {
t.Fatalf("settings writes = %d, want selected values written before compensation", fake.writes)
}
if fake.settingsExist != prior.exists || string(fake.settingsRaw) != string(prior.raw) {
t.Fatalf("restored exists/raw = %t/%q, want %t/%q", fake.settingsExist, fake.settingsRaw, prior.exists, prior.raw)
}
if fake.settingsReads < 3 {
t.Fatalf("settings reads = %d, want prior effective state, requested readback, and restored default verification", fake.settingsReads)
}
})
}
}
// Catches tht reading the legacy public model route instead of the admin-only closed Pi
// Management choices before it writes shared installation defaults.
func TestConfigureLoadsDedicatedClosedOptionsWritesRealCoreSettingsAndUsesUpstreamIdentity(t *testing.T) {
fake := newFakeRunner()
if err := Configure(context.Background(), fake, Defaults{Provider: "provider", Model: "model", Thinking: "medium"}); err != nil {
t.Fatal(err)
}
assertCalled(t, fake.calls, "/pi-management/options")
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: tht-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, Defaults{Provider: "provider", Model: "unknown", Thinking: "medium"}); err == nil {
t.Fatal("expected unknown model rejection")
}
}
// Catches a smoke check that composes health/models/settings itself and drifts from the dedicated
// backend contract, rather than retaining only the independent in-container version signal.
func TestTestUsesDedicatedSmokeEndpointAndIndependentImageVersionProbe(t *testing.T) {
fake := newFakeRunner()
if err := Test(context.Background(), fake); err != nil {
t.Fatalf("Test() error = %v", err)
}
for _, command := range []string{"pi --version", "/pi-management/test", "x-thoth-principal-subject: tht-maintenance"} {
assertCalled(t, fake.calls, command)
}
for _, legacy := range []string{"/health", "/models", "/settings"} {
if strings.Contains(strings.Join(fake.calls, "\n"), legacy) {
t.Fatalf("Pi smoke invoked legacy endpoint %q: %s", legacy, strings.Join(fake.calls, "\n"))
}
}
if got := strings.Join(fake.calls, "\n"); strings.Contains(got, "secret") {
t.Fatalf("probe commands expose secret: %s", got)
}
}
// Catches an ignored negative ready result from the backend smoke endpoint, which would report a
// successfully verified candidate image while its configured Pi runtime is unusable.
func TestTestRequiresDedicatedSmokeEndpointToReportReady(t *testing.T) {
fake := newFakeRunner()
fake.piManagementTestWire = `{"ready":false,"message":"provider unavailable"}`
if err := Test(context.Background(), fake); err == nil || !strings.Contains(err.Error(), "Pi smoke response is not ready") {
t.Fatalf("Test() error = %v, want negative dedicated smoke result", err)
}
fake.piManagementTestWire = `{"ready":true}`
if err := Test(context.Background(), fake); err != nil {
t.Fatalf("Test() exact match error = %v", err)
}
assertCalled(t, fake.calls, "pi --version")
}
// Catches tht accepting a reasoning level that the backend did not publish as a closed
// installation option, which would bypass the Pi Management validation surface.
func TestConfigureRejectsReasoningOutsideDedicatedClosedOptions(t *testing.T) {
fake := newFakeRunner()
fake.piManagementOptionsWire = `{"providers":["provider"],"models":[{"provider":"provider","id":"model"}],"reasoning":["low"]}`
if err := Configure(context.Background(), fake, Defaults{Provider: "provider", Model: "model", Thinking: "high"}); err == nil || !strings.Contains(err.Error(), "Pi options") {
t.Fatalf("Configure() error = %v, want closed reasoning rejection", err)
}
}
+35
View File
@@ -0,0 +1,35 @@
//go:build !windows
package pi
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
}
dir, err := os.Open(directory)
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()
}
+31
View File
@@ -0,0 +1,31 @@
//go:build windows
package pi
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
}
to, err := windows.UTF16PtrFromString(target)
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
}
+24
View File
@@ -0,0 +1,24 @@
package pi
// RecoveryRequiredError marks a result whose immediate state may be safe but whose durability
// was explicitly not acknowledged. Callers must not report success or clear maintenance.
type RecoveryRequiredError struct {
Operation string
Cause error
}
func (e *RecoveryRequiredError) Error() string {
return e.Operation + ": recovery required"
}
func (e *RecoveryRequiredError) Unwrap() error {
return e.Cause
}
func (e *RecoveryRequiredError) RecoveryRequired() bool {
return true
}
func recoveryRequired(operation string, cause error) error {
return &RecoveryRequiredError{Operation: operation, Cause: cause}
}
+360
View File
@@ -0,0 +1,360 @@
package pi
import (
"context"
"errors"
"fmt"
"os"
"path/filepath"
)
var (
ErrInterruptedRestart = errors.New("a previous Pi restart is incomplete; recover lifecycle maintenance before another Pi lifecycle operation")
errRestartImageDrift = errors.New("core image changed during Pi restart")
errRestartConfigurationDrift = errors.New("external endpoint configuration changed during Pi restart")
errRestartMountDrift = errors.New("core persistence mount contract changed during Pi restart")
errRestartConfirmation = &restartDiagnosticError{
message: "restart requires --yes after reviewing the planned Pi core recreation",
cause: ErrConfirmationRequired,
}
errRestartActiveSessions = &restartDiagnosticError{
message: "active sessions must be drained before restarting Pi; use --drain only after they are complete",
cause: ErrActiveSessions,
}
)
type restartDiagnosticError struct {
message string
cause error
}
func (e *restartDiagnosticError) Error() string { return e.message }
func (e *restartDiagnosticError) Unwrap() error { return e.cause }
type RestartRequest struct {
StatePath string
UpdateStatePath string
Confirm bool
Drain bool
}
type RestartResult struct {
StatePath string
Version string
}
func Restart(ctx context.Context, runner Runner, request RestartRequest) (RestartResult, error) {
return restartWithHooks(ctx, runner, request, defaultLifecycleHooks)
}
func restartWithHooks(
ctx context.Context,
runner Runner,
request RestartRequest,
hooks lifecycleHooks,
) (result RestartResult, retErr error) {
if err := validateRestartStatePaths(request.StatePath, request.UpdateStatePath); err != nil {
return RestartResult{}, err
}
lock, err := acquireLock(request.StatePath)
if err != nil {
return RestartResult{StatePath: request.StatePath}, err
}
defer lock.Release()
if !request.Confirm {
return RestartResult{StatePath: request.StatePath}, errRestartConfirmation
}
if state, err := readState(request.UpdateStatePath); err == nil && stateNeedsRecovery(state) {
return RestartResult{StatePath: request.StatePath}, ErrInterruptedUpdate
} else if err != nil && !errors.Is(err, os.ErrNotExist) {
return RestartResult{StatePath: request.StatePath}, err
}
if err := prepareLifecycleMutation(request.StatePath, hooks.removeFile); err != nil {
return RestartResult{StatePath: request.StatePath}, err
}
clearMaintenance := true
mutationStarted := false
var state State
defer func() {
if !clearMaintenance {
return
}
if clearErr := setMaintenance(context.Background(), runner, false); clearErr != nil {
result.StatePath = request.StatePath
if mutationStarted {
clearMaintenance = false
cause := clearErr
if writeErr := hooks.writeState(request.StatePath, state); writeErr != nil {
cause = errors.Join(cause, fmt.Errorf("restart recovery state could not be restored: %w", writeErr))
}
retErr = errors.Join(retErr, recoveryRequired("Pi restart maintenance cleanup failed", cause))
return
}
retErr = errors.Join(retErr, fmt.Errorf("maintenance admission gate could not be cleared: %w", clearErr))
}
}()
if err := setMaintenance(ctx, runner, true); err != nil {
return RestartResult{StatePath: request.StatePath}, err
}
if err := waitForInactiveSessions(ctx, runner, request.Drain, hooks.sleep); err != nil {
return RestartResult{StatePath: request.StatePath}, restartDiagnostic(err)
}
if err := Doctor(ctx, runner); err != nil {
return RestartResult{StatePath: request.StatePath}, err
}
version, err := Status(ctx, runner)
if err != nil {
return RestartResult{StatePath: request.StatePath}, err
}
configured, err := renderedCore(ctx, runner)
if err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, err
}
previous, err := runningImage(ctx, runner, configured.Reference)
if err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, err
}
previous.ConfigurationSHA = configured.ConfigurationSHA
transaction := lifecycleTransaction(request.StatePath)
previous.Reference = lifecycleImageTag(transaction, "restart")
if err := tagImage(ctx, runner, previous.ID, previous.Reference, "restart image pin"); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, err
}
state = State{
Transaction: transaction,
Phase: PhasePreflight,
Target: Target{Version: version, Source: "restart"},
Previous: previous,
}
if err := hooks.writeState(request.StatePath, state); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, err
}
overridePath := lifecycleOverridePath(request.StatePath, transaction)
if err := writeLifecycleOverride(overridePath, previous.Reference); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, err
}
lifecycle := composeOverrideRunner{Runner: runner, path: overridePath}
if running, err := activeSessions(ctx, runner); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, err
} else if running {
return RestartResult{StatePath: request.StatePath, Version: version}, errRestartActiveSessions
}
state.MutationStarted = true
if err := hooks.writeState(request.StatePath, state); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, err
}
mutationStarted = true
clearMaintenance = false
if err := recreateCoreWithoutImageChanges(ctx, lifecycle); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, recoveryRequired("Pi restart core recreation failed", err)
}
if err := ensureMaintenance(ctx, lifecycle); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, recoveryRequired("Pi restart maintenance proof failed", err)
}
state.Phase = PhaseRecreated
if err := hooks.writeState(request.StatePath, state); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, recoveryRequired("Pi restart recreation state could not be recorded", err)
}
if err := verifyRestart(ctx, lifecycle, version, previous); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, recoveryRequired("Pi restart verification failed", err)
}
state.Phase = PhaseVerified
if err := hooks.writeState(request.StatePath, state); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, recoveryRequired("Pi restart verification state could not be recorded", err)
}
if err := hooks.removeFile(overridePath); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, recoveryRequired("Pi restart image override could not be removed", err)
}
if err := hooks.removeFile(request.StatePath); err != nil {
return RestartResult{StatePath: request.StatePath, Version: version}, recoveryRequired("Pi restart recovery state could not be removed", err)
}
clearMaintenance = true
return RestartResult{StatePath: request.StatePath, Version: version}, nil
}
func recreateCoreWithoutImageChanges(ctx context.Context, runner Runner) error {
result, err := runCompose(
ctx,
runner,
"up",
"--detach",
"--wait",
"--wait-timeout",
"45",
"--no-deps",
"--force-recreate",
"--no-build",
"--pull",
"never",
"core",
)
if err != nil {
return commandError("core recreation", result, err)
}
return nil
}
func verifyRestart(ctx context.Context, runner Runner, wanted string, previous Image) error {
if err := Doctor(ctx, runner); err != nil {
return err
}
version, err := Status(ctx, runner)
if err != nil {
return err
}
if version != wanted {
return errors.New("Pi version changed during core restart")
}
configured, err := renderedCore(ctx, runner)
if err != nil {
return err
}
after, err := runningImage(ctx, runner, configured.Reference)
if err != nil {
return err
}
if after.ID != previous.ID {
return errRestartImageDrift
}
if configured.ConfigurationSHA != previous.ConfigurationSHA {
return errRestartConfigurationDrift
}
if !sameMounts(previous.Mounts, after.Mounts) {
return errRestartMountDrift
}
return nil
}
func RecoverLifecycleMaintenance(
ctx context.Context,
runner Runner,
updateStatePath string,
restartStatePath string,
confirm bool,
) error {
if err := validateRestartStatePaths(restartStatePath, updateStatePath); err != nil {
return err
}
if !confirm {
return errRestartConfirmation
}
lock, err := acquireLock(restartStatePath)
if err != nil {
return err
}
defer lock.Release()
restartState, restartErr := readState(restartStatePath)
if restartErr == nil {
if err := validateRestartRecoveryState(restartState); err != nil {
return err
}
restartOverride := lifecycleOverridePath(restartStatePath, restartState.Transaction)
if restartState.MutationStarted {
if err := tagImage(ctx, runner, restartState.Previous.ID, restartState.Previous.Reference, "restart recovery image pin"); err != nil {
return recoveryRequired("Pi restart recovery image pin could not be restored", err)
}
if err := writeLifecycleOverride(restartOverride, restartState.Previous.Reference); err != nil {
return recoveryRequired("Pi restart recovery image override could not be restored", err)
}
lifecycle := composeOverrideRunner{Runner: runner, path: restartOverride}
if err := ensureMaintenance(ctx, lifecycle); err != nil {
return recoveryRequired("Pi restart maintenance recovery failed", err)
}
if err := verifyRestart(ctx, lifecycle, restartState.Target.Version, restartState.Previous); err != nil {
return recoveryRequired("Pi restart recovery verification failed", err)
}
}
if err := durableRemove(restartOverride); err != nil {
return recoveryRequired("Pi restart recovery image override could not be removed", err)
}
if err := durableRemove(restartStatePath); err != nil {
return recoveryRequired("Pi restart recovery state could not be removed", err)
}
} else if !errors.Is(restartErr, os.ErrNotExist) {
return restartErr
}
return recoverMaintenanceLocked(ctx, runner, updateStatePath)
}
func restartDiagnostic(err error) error {
if errors.Is(err, ErrActiveSessions) {
return errRestartActiveSessions
}
return err
}
func validateRestartRecoveryState(state State) error {
invalid := func(reason string) error {
return fmt.Errorf("%w: restart recovery state is invalid: %s", ErrInvalidRequest, reason)
}
if state.Transaction == "" {
return invalid("transaction is missing")
}
if state.Target.Source != "restart" || state.Target.Version == "" {
return invalid("target is not a restart with a recorded version")
}
if state.Previous.ConfigurationSHA == "" {
return invalid("previous external configuration identity is missing")
}
if state.Candidate.ID != "" || state.Candidate.Reference != "" || len(state.Candidate.Mounts) != 0 ||
state.Candidate.MountFingerprint != "" || state.Candidate.ConfigurationSHA != "" {
return invalid("restart state contains image candidate metadata")
}
switch state.Phase {
case PhasePreflight:
return nil
case PhaseRecreated, PhaseVerified:
if state.MutationStarted {
return nil
}
return invalid("post-recreation phase has no mutation marker")
default:
return invalid("phase is not valid for restart")
}
}
func validateRestartStatePaths(restartStatePath, updateStatePath string) error {
if restartStatePath == "" {
return errors.New("restart state path is required")
}
if updateStatePath == "" {
return errors.New("update state path is required")
}
if filepath.Clean(restartStatePath) == filepath.Clean(updateStatePath) {
return fmt.Errorf("%w: restart and update state paths must remain separate", ErrInvalidRequest)
}
if lifecycleLockPath(restartStatePath) != lifecycleLockPath(updateStatePath) {
return fmt.Errorf("%w: restart and update state paths must share one lifecycle control directory", ErrInvalidRequest)
}
return nil
}
func pairedRestartStatePath(updateStatePath string) string {
return filepath.Join(filepath.Dir(updateStatePath), "restart-state.json")
}
func prepareLifecycleMutation(restartStatePath string, removeFile func(string) error) error {
state, err := readState(restartStatePath)
if errors.Is(err, os.ErrNotExist) {
return nil
}
if err != nil {
return fmt.Errorf("restart recovery state could not be validated: %w", err)
}
if err := validateRestartRecoveryState(state); err != nil {
return err
}
if state.Phase != PhaseVerified || !state.MutationStarted {
return ErrInterruptedRestart
}
if err := removeFile(lifecycleOverridePath(restartStatePath, state.Transaction)); err != nil {
return recoveryRequired("verified restart override could not be cleaned up", err)
}
if err := removeFile(restartStatePath); err != nil {
return recoveryRequired("verified restart state could not be cleaned up", err)
}
return nil
}
+599
View File
@@ -0,0 +1,599 @@
package pi
import (
"context"
"errors"
"os"
"path/filepath"
"strings"
"testing"
"time"
)
func TestRestartRequiresConfirmationWithoutInvokingCompose(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
_, err := Restart(context.Background(), fake, RestartRequest{
StatePath: filepath.Join(dir, "restart-state.json"),
UpdateStatePath: filepath.Join(dir, "update-state.json"),
})
if !errors.Is(err, ErrConfirmationRequired) {
t.Fatalf("Restart() error = %v, want ErrConfirmationRequired", err)
}
if got, want := err.Error(), "restart requires --yes after reviewing the planned Pi core recreation"; got != want {
t.Fatalf("Restart() error text = %q, want %q", got, want)
}
assertNotCalled(t, fake.calls, "compose")
}
func TestRestartDrainsRecreatesOnlyCoreAndRetainsImage(t *testing.T) {
fake := newFakeRunner()
fake.sessionsWire = `[{"status":"open","archived":false}]`
dir := t.TempDir()
hooks := defaultLifecycleHooks
sleepCalls := 0
hooks.sleep = func(duration time.Duration) {
sleepCalls++
if duration != time.Second {
t.Fatalf("drain sleep = %s, want %s", duration, time.Second)
}
fake.sessionsWire = `[]`
}
result, err := restartWithHooks(context.Background(), fake, RestartRequest{
StatePath: filepath.Join(dir, "restart-state.json"),
UpdateStatePath: filepath.Join(dir, "update-state.json"),
Confirm: true,
Drain: true,
}, hooks)
if err != nil {
t.Fatal(err)
}
if result.Version != fake.version {
t.Fatalf("version = %q, want %q", result.Version, fake.version)
}
if sleepCalls != 1 {
t.Fatalf("drain sleep calls = %d, want 1", sleepCalls)
}
assertCalled(t, fake.calls, "up --detach --wait --wait-timeout 45 --no-deps --force-recreate --no-build --pull never core")
assertNotCalled(t, fake.calls, "compose build --pull")
for _, call := range fake.calls {
if strings.HasPrefix(call, "pull ") {
t.Fatalf("restart invoked direct image pull: %s", call)
}
}
assertNotCalled(t, fake.calls, "frontend")
if _, err := os.Stat(result.StatePath); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("successful restart state still exists: %v", err)
}
}
func TestRestartPinsCapturedImageWhenConfiguredTagMovesBeforeRecreate(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
hooks := defaultLifecycleHooks
write := hooks.writeState
hooks.writeState = func(path string, state State) error {
if state.Phase == PhasePreflight && state.MutationStarted {
fake.tags[fake.configuredImage] = "sha256:moved-configured-tag"
fake.imageVersions["sha256:moved-configured-tag"] = "9.99.0"
}
return write(path, state)
}
restartStatePath := filepath.Join(dir, "restart-state.json")
result, err := restartWithHooks(context.Background(), fake, RestartRequest{
StatePath: restartStatePath,
UpdateStatePath: filepath.Join(dir, "update-state.json"),
Confirm: true,
}, hooks)
if err != nil {
t.Fatalf("restartWithHooks() error = %v", err)
}
if result.Version != "0.80.3" || fake.currentImage != "sha256:old" {
t.Fatalf("restart result=%+v image=%q; want captured 0.80.3 / sha256:old", result, fake.currentImage)
}
assertCalled(t, fake.calls, "image tag sha256:old thothii-core:tht-")
assertCalled(t, fake.calls, "pi-lifecycle-")
if matches, globErr := filepath.Glob(filepath.Join(dir, "pi-lifecycle-*.yaml")); globErr != nil || len(matches) != 0 {
t.Fatalf("successful restart overrides = %v, error = %v; want safe cleanup", matches, globErr)
}
}
func TestRestartRefusesActiveSessionsWithoutDrain(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
fake.activeSessions = true
_, err := Restart(context.Background(), fake, RestartRequest{
StatePath: filepath.Join(dir, "restart-state.json"),
UpdateStatePath: filepath.Join(dir, "update-state.json"),
Confirm: true,
})
if !errors.Is(err, ErrActiveSessions) {
t.Fatalf("Restart() error = %v, want ErrActiveSessions", err)
}
if got, want := err.Error(), "active sessions must be drained before restarting Pi; use --drain only after they are complete"; got != want {
t.Fatalf("Restart() error text = %q, want %q", got, want)
}
if fake.maintenance {
t.Fatal("maintenance remained active after refusing pre-mutation restart")
}
if fake.recreated {
t.Fatal("core was recreated with active sessions")
}
}
func TestRestartActivationFailureClearsPreMutationMaintenance(t *testing.T) {
for _, failure := range []string{
"maintenance-activate-durability",
"maintenance-activate-durability-without-status-flag",
} {
t.Run(failure, func(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
fake.fail = failure
_, err := Restart(context.Background(), fake, RestartRequest{
StatePath: filepath.Join(dir, "restart-state.json"),
UpdateStatePath: filepath.Join(dir, "update-state.json"),
Confirm: true,
})
var recovery *RecoveryRequiredError
if !errors.As(err, &recovery) {
t.Fatalf("Restart() error = %v, want RecoveryRequiredError", err)
}
if fake.maintenance {
t.Fatal("pre-mutation activation failure left maintenance active")
}
if fake.recreated {
t.Fatal("pre-mutation activation failure recreated core")
}
assertCalled(t, fake.calls, "/internal/maintenance/activate")
assertCalled(t, fake.calls, "/internal/maintenance/deactivate")
})
}
}
func TestRestartRefusesInterruptedUpdateOrRestartState(t *testing.T) {
for _, stateFile := range []string{"update-state.json", "restart-state.json"} {
t.Run(stateFile, func(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
previous := stateImageForTest(t, fake)
writeStateForTest(t, filepath.Join(dir, stateFile), State{
Transaction: "interrupted",
Phase: PhaseRecreated,
Target: Target{Version: fake.version, Source: "restart"},
Previous: previous,
MutationStarted: true,
})
fake.calls = nil
_, err := Restart(context.Background(), fake, RestartRequest{
StatePath: filepath.Join(dir, "restart-state.json"),
UpdateStatePath: filepath.Join(dir, "update-state.json"),
Confirm: true,
})
if err == nil {
t.Fatal("Restart() accepted interrupted lifecycle state")
}
assertNotCalled(t, fake.calls, "compose")
})
}
}
func TestRestartPreflightFailureNeverRecreatesCoreAndClearsMaintenance(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
fake.fail = "preflight"
_, err := Restart(context.Background(), fake, RestartRequest{
StatePath: filepath.Join(dir, "restart-state.json"),
UpdateStatePath: filepath.Join(dir, "update-state.json"),
Confirm: true,
})
if err == nil {
t.Fatal("Restart() error = nil, want preflight failure")
}
if fake.recreated || fake.currentImage != "sha256:old" {
t.Fatalf("preflight failure mutated core: recreated=%t image=%q", fake.recreated, fake.currentImage)
}
if fake.maintenance {
t.Fatal("maintenance remained active after preflight failure")
}
}
func TestRestartPostRecreateFailureKeepsMaintenanceAndRecoveryState(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
fake.fail = "health"
statePath := filepath.Join(dir, "restart-state.json")
_, err := Restart(context.Background(), fake, RestartRequest{
StatePath: statePath,
UpdateStatePath: filepath.Join(dir, "update-state.json"),
Confirm: true,
})
var recovery *RecoveryRequiredError
if !errors.As(err, &recovery) {
t.Fatalf("Restart() error = %v, want RecoveryRequiredError", err)
}
if !fake.recreated {
t.Fatal("post-recreate failure occurred before core recreation")
}
if !fake.maintenance {
t.Fatal("maintenance was cleared after post-recreate failure")
}
state, stateErr := readState(statePath)
if stateErr != nil || !state.MutationStarted {
t.Fatalf("restart recovery state = %+v, %v; want durable mutation state", state, stateErr)
}
overridePath := lifecycleOverridePath(statePath, state.Transaction)
selected, overrideErr := readLifecycleOverride(overridePath)
if overrideErr != nil || selected != state.Previous.Reference || fake.tags[selected] != state.Previous.ID {
t.Fatalf("restart override = %q, %v; want retained exact image %q", selected, overrideErr, state.Previous.ID)
}
}
func TestRestartMaintenanceClearFailureRestoresRecoveryState(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
fake.fail = "maintenance-clear"
statePath := filepath.Join(dir, "restart-state.json")
_, err := Restart(context.Background(), fake, RestartRequest{
StatePath: statePath,
UpdateStatePath: filepath.Join(dir, "update-state.json"),
Confirm: true,
})
var recovery *RecoveryRequiredError
if !errors.As(err, &recovery) {
t.Fatalf("Restart() error = %v, want RecoveryRequiredError", err)
}
if !fake.maintenance {
t.Fatal("maintenance was cleared despite deactivation failure")
}
state, stateErr := readState(statePath)
if stateErr != nil || state.Phase != PhaseVerified || !state.MutationStarted {
t.Fatalf("restart recovery state = %+v, %v; want durable verified mutation state", state, stateErr)
}
}
func TestRecoverLifecycleMaintenanceVerifiesAndClearsRestartState(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
fake.maintenance = true
previous := stateImageForTest(t, fake)
restartStatePath := filepath.Join(dir, "restart-state.json")
updateStatePath := filepath.Join(dir, "update-state.json")
writeStateForTest(t, restartStatePath, State{
Transaction: "restart-recovery",
Phase: PhaseRecreated,
Target: Target{Version: fake.version, Source: "restart"},
Previous: previous,
MutationStarted: true,
})
restartState, err := readState(restartStatePath)
if err != nil {
t.Fatal(err)
}
restartOverride := lifecycleOverridePath(restartStatePath, restartState.Transaction)
if err := writeLifecycleOverride(restartOverride, restartState.Previous.Reference); err != nil {
t.Fatal(err)
}
delete(fake.tags, restartState.Previous.Reference)
fake.tags[fake.configuredImage] = "sha256:moved-before-recovery"
candidate := previous
candidate.Reference = "thothii-core:tht-recover-candidate"
writeStateForTest(t, updateStatePath, State{
Transaction: "update-recovery",
Phase: PhasePromoting,
Target: Target{Version: fake.version, Source: string(BuildSource)},
Previous: previous,
Candidate: candidate,
MutationStarted: true,
})
if err := writeLifecycleOverride(currentImageOverridePath(updateStatePath), candidate.Reference); err != nil {
t.Fatal(err)
}
fake.calls = nil
if err := RecoverLifecycleMaintenance(context.Background(), fake, updateStatePath, restartStatePath, true); err != nil {
t.Fatalf("RecoverLifecycleMaintenance() error = %v", err)
}
if _, err := os.Stat(restartStatePath); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("restart recovery state still exists: %v", err)
}
if _, err := os.Stat(restartOverride); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("restart recovery override still exists: %v", err)
}
if fake.tags[restartState.Previous.Reference] != restartState.Previous.ID {
t.Fatalf("restart recovery pin = %q, want %q", fake.tags[restartState.Previous.Reference], restartState.Previous.ID)
}
assertCalled(t, fake.calls, restartOverride)
if fake.maintenance {
t.Fatal("maintenance remained active after both lifecycle states were verified")
}
updateState, err := readState(updateStatePath)
if err != nil || updateState.Phase != PhaseVerified {
t.Fatalf("update recovery state = %+v, %v; want verified image rollback metadata", updateState, err)
}
deactivate := callIndex(fake.calls, "/internal/maintenance/deactivate")
lastVerification := lastCallIndexBefore(fake.calls, "/pi-management/test", deactivate)
if deactivate < 0 || lastVerification < 0 {
t.Fatalf("calls = %v; want verification before maintenance deactivation", fake.calls)
}
verificationCount := 0
for index := 0; index < deactivate; index++ {
if strings.Contains(fake.calls[index], "/internal/maintenance/deactivate") {
t.Fatalf("maintenance reopened before combined verification: %v", fake.calls)
}
if strings.Contains(fake.calls[index], "/pi-management/test") {
verificationCount++
}
}
if verificationCount < 3 {
t.Fatalf("verification calls before maintenance deactivation = %d, want restart, update, and final proofs: %v", verificationCount, fake.calls)
}
}
func TestRecoverLifecycleMaintenanceRejectsMalformedRestartState(t *testing.T) {
for _, test := range []struct {
name string
phase Phase
source string
}{
{name: "recreated_without_mutation_marker", phase: PhaseRecreated, source: "restart"},
{name: "non_restart_source", phase: PhasePreflight, source: string(BuildSource)},
} {
t.Run(test.name, func(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
fake.maintenance = true
restartStatePath := filepath.Join(dir, "restart-state.json")
updateStatePath := filepath.Join(dir, "update-state.json")
writeStateForTest(t, restartStatePath, State{
Transaction: "malformed-restart",
Phase: test.phase,
Target: Target{Version: fake.version, Source: test.source},
Previous: stateImageForTest(t, fake),
})
fake.calls = nil
err := RecoverLifecycleMaintenance(
context.Background(), fake, updateStatePath, restartStatePath, true,
)
if !errors.Is(err, ErrInvalidRequest) || !strings.Contains(err.Error(), "restart recovery state is invalid") {
t.Fatalf("RecoverLifecycleMaintenance() error = %v, want invalid restart recovery state", err)
}
if !fake.maintenance {
t.Fatal("malformed restart state reopened admission")
}
if _, stateErr := os.Stat(restartStatePath); stateErr != nil {
t.Fatalf("malformed restart state was removed: %v", stateErr)
}
assertNotCalled(t, fake.calls, "/internal/maintenance/deactivate")
})
}
}
func TestRestartRefusesMalformedRestartStateWithoutInvokingCompose(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
restartStatePath := filepath.Join(dir, "restart-state.json")
writeStateForTest(t, restartStatePath, State{
Transaction: "malformed-restart",
Phase: PhasePreflight,
Target: Target{Version: fake.version, Source: string(BuildSource)},
Previous: stateImageForTest(t, fake),
})
fake.calls = nil
_, err := Restart(context.Background(), fake, RestartRequest{
StatePath: restartStatePath,
UpdateStatePath: filepath.Join(dir, "update-state.json"),
Confirm: true,
})
if !errors.Is(err, ErrInvalidRequest) || !strings.Contains(err.Error(), "restart recovery state is invalid") {
t.Fatalf("Restart() error = %v, want invalid restart recovery state", err)
}
if _, stateErr := os.Stat(restartStatePath); stateErr != nil {
t.Fatalf("malformed restart state was removed: %v", stateErr)
}
assertNotCalled(t, fake.calls, "compose")
}
func TestRestartDurabilityFailureBoundaries(t *testing.T) {
injected := errors.New("injected restart durability failure")
for _, test := range []struct {
name string
configure func(*fakeRunner, *lifecycleHooks)
wantPhase Phase
wantMutation bool
wantRecreated bool
wantMaintenance bool
wantRecovery bool
wantInjected bool
}{
{
name: "mutation_marker_write",
configure: func(_ *fakeRunner, hooks *lifecycleHooks) {
write := hooks.writeState
hooks.writeState = func(path string, state State) error {
if state.Phase == PhasePreflight && state.MutationStarted {
return injected
}
return write(path, state)
}
},
wantPhase: PhasePreflight,
wantInjected: true,
},
{
name: "core_recreation",
configure: func(fake *fakeRunner, _ *lifecycleHooks) {
fake.fail = "recreate"
},
wantPhase: PhasePreflight,
wantMutation: true,
wantRecreated: true,
wantMaintenance: true,
wantRecovery: true,
},
{
name: "maintenance_proof",
configure: func(fake *fakeRunner, _ *lifecycleHooks) {
fake.fail = "maintenance-proof"
},
wantPhase: PhasePreflight,
wantMutation: true,
wantRecreated: true,
wantMaintenance: true,
wantRecovery: true,
},
{
name: "recreated_phase_write",
configure: func(_ *fakeRunner, hooks *lifecycleHooks) {
write := hooks.writeState
hooks.writeState = func(path string, state State) error {
if state.Phase == PhaseRecreated {
return injected
}
return write(path, state)
}
},
wantPhase: PhasePreflight,
wantMutation: true,
wantRecreated: true,
wantMaintenance: true,
wantRecovery: true,
wantInjected: true,
},
{
name: "verified_phase_write",
configure: func(_ *fakeRunner, hooks *lifecycleHooks) {
write := hooks.writeState
hooks.writeState = func(path string, state State) error {
if state.Phase == PhaseVerified {
return injected
}
return write(path, state)
}
},
wantPhase: PhaseRecreated,
wantMutation: true,
wantRecreated: true,
wantMaintenance: true,
wantRecovery: true,
wantInjected: true,
},
{
name: "restart_state_removal",
configure: func(_ *fakeRunner, hooks *lifecycleHooks) {
hooks.removeFile = func(string) error { return injected }
},
wantPhase: PhaseVerified,
wantMutation: true,
wantRecreated: true,
wantMaintenance: true,
wantRecovery: true,
wantInjected: true,
},
} {
t.Run(test.name, func(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
hooks := defaultLifecycleHooks
test.configure(fake, &hooks)
statePath := filepath.Join(dir, "restart-state.json")
_, err := restartWithHooks(context.Background(), fake, RestartRequest{
StatePath: statePath,
UpdateStatePath: filepath.Join(dir, "update-state.json"),
Confirm: true,
}, hooks)
if err == nil {
t.Fatal("restartWithHooks() error = nil, want injected boundary failure")
}
var recovery *RecoveryRequiredError
if got := errors.As(err, &recovery); got != test.wantRecovery {
t.Fatalf("recovery-required = %t, want %t; error = %v", got, test.wantRecovery, err)
}
if test.wantInjected && !errors.Is(err, injected) {
t.Fatalf("restartWithHooks() error = %v, want injected durability cause", err)
}
state, stateErr := readState(statePath)
if stateErr != nil {
t.Fatalf("readState() error = %v", stateErr)
}
if state.Phase != test.wantPhase || state.MutationStarted != test.wantMutation {
t.Fatalf("state = %+v, want phase=%q mutation=%t", state, test.wantPhase, test.wantMutation)
}
if fake.recreated != test.wantRecreated {
t.Fatalf("recreated = %t, want %t", fake.recreated, test.wantRecreated)
}
if fake.maintenance != test.wantMaintenance {
t.Fatalf("maintenance = %t, want %t", fake.maintenance, test.wantMaintenance)
}
})
}
}
func TestRestartRejectsImageConfigurationAndMountDrift(t *testing.T) {
for _, test := range []struct {
failure string
want error
}{
{failure: "image-drift", want: errRestartImageDrift},
{failure: "config-drift", want: errRestartConfigurationDrift},
{failure: "mount-drift", want: errRestartMountDrift},
} {
t.Run(test.failure, func(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
fake.fail = test.failure
statePath := filepath.Join(dir, "restart-state.json")
_, err := Restart(context.Background(), fake, RestartRequest{
StatePath: statePath,
UpdateStatePath: filepath.Join(dir, "update-state.json"),
Confirm: true,
})
if !errors.Is(err, test.want) {
t.Fatalf("Restart() error = %v, want errors.Is(..., %v)", err, test.want)
}
if !fake.maintenance {
t.Fatal("maintenance was cleared after restart identity drift")
}
if _, stateErr := os.Stat(statePath); stateErr != nil {
t.Fatalf("restart recovery state missing after drift: %v", stateErr)
}
})
}
}
func TestRestartPreservesVerifiedUpdateState(t *testing.T) {
dir := t.TempDir()
fake := newFakeRunner()
updateStatePath := filepath.Join(dir, "update-state.json")
writeStateForTest(t, updateStatePath, State{
Transaction: "verified-update",
Phase: PhaseVerified,
Target: Target{Version: fake.version, Source: string(BuildSource)},
Previous: stateImageForTest(t, fake),
})
before := readStateBytes(t, updateStatePath)
if _, err := Restart(context.Background(), fake, RestartRequest{
StatePath: filepath.Join(dir, "restart-state.json"),
UpdateStatePath: updateStatePath,
Confirm: true,
}); err != nil {
t.Fatal(err)
}
after := readStateBytes(t, updateStatePath)
if string(after) != string(before) {
t.Fatal("restart changed verified update rollback metadata")
}
}
+223
View File
@@ -0,0 +1,223 @@
// Package pi implements host-side lifecycle operations for the Pi bundled in core.
package pi
import (
"crypto/sha256"
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"time"
"github.com/gofrs/flock"
)
const stateFileVersion = 4
// Phase describes the durable point reached by a Pi update.
type Phase string
const (
PhasePreflight Phase = "preflight"
PhaseBuilding Phase = "building"
PhaseRecreated Phase = "recreated"
PhasePromoting Phase = "promoting"
PhaseVerified Phase = "verified"
PhaseRolledBack Phase = "rolled_back"
PhaseFailed Phase = "failed"
PhaseNoop Phase = "noop"
)
// 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"`
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"`
SourceSHA256 string `json:"source_sha256"`
SourceAliases []string `json:"-"`
Destination string `json:"destination"`
RW bool `json:"rw"`
Options string `json:"options,omitempty"`
}
// Target records the immutable input selected by the operator. Source is build, restart, or a
// digest-pinned image reference; it intentionally never contains credentials.
type Target struct {
Version string `json:"version"`
Source string `json:"source"`
}
// 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"`
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"`
MutationStarted bool `json:"mutation_started,omitempty"`
Error string `json:"error,omitempty"`
}
func readState(path string) (State, error) {
contents, err := os.ReadFile(path)
if err != nil {
return State{}, err
}
var state State
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 == "" || state.Previous.MountFingerprint == "" {
return State{}, errors.New("update recovery state is incomplete")
}
if mountFingerprint(state.Previous.Mounts) != state.Previous.MountFingerprint || (state.Candidate.ID != "" && mountFingerprint(state.Candidate.Mounts) != state.Candidate.MountFingerprint) {
return State{}, errors.New("update recovery state mount fingerprint is invalid")
}
return state, nil
}
func writeState(path string, state State) error {
if state.Previous.ID == "" || state.Previous.Reference == "" || state.Previous.MountFingerprint == "" {
return errors.New("refusing to write incomplete update recovery state")
}
state.Version = stateFileVersion
state.UpdatedAt = time.Now().UTC()
contents, err := json.MarshalIndent(state, "", " ")
if err != nil {
return fmt.Errorf("encode update recovery state: %w", err)
}
contents = append(contents, '\n')
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 mountSourceAliases(mountType, source, goos string) []string {
if mountType != "bind" || goos != "darwin" {
return nil
}
source = filepath.Clean(source)
var alias string
switch {
case strings.HasPrefix(source, "/host_mnt/private/var/"), strings.HasPrefix(source, "/host_mnt/Users/"):
alias = strings.TrimPrefix(source, "/host_mnt")
case strings.HasPrefix(source, "/private/var/"), strings.HasPrefix(source, "/Users/"):
alias = "/host_mnt" + source
default:
return nil
}
return []string{mountSourceHash(alias)}
}
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 {
file *flock.Flock
metadata string
}
var ErrLockHeld = errors.New("another Pi update, restart, or rollback is already in progress")
func lifecycleLockPath(statePath string) string {
return filepath.Join(filepath.Dir(statePath), "pi-lifecycle.lock")
}
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")
}
path := lifecycleLockPath(statePath)
file := flock.New(path, flock.SetPermissions(0o600))
locked, err := file.TryLock()
if err != nil {
return nil, errors.New("could not acquire Pi update lock")
}
if !locked {
return nil, ErrLockHeld
}
host, err := os.Hostname()
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.Unlock()
return nil, errors.New("could not record Pi update lock owner")
}
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() {
_ = durableRemove(l.metadata)
_ = l.file.Unlock()
}
+76
View File
@@ -0,0 +1,76 @@
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 TestUpdateAndRestartStatePathsShareOneLifecycleLock(t *testing.T) {
dir := t.TempDir()
first, err := acquireLock(filepath.Join(dir, "update-state.json"))
if err != nil {
t.Fatal(err)
}
defer first.Release()
second, err := acquireLock(filepath.Join(dir, "restart-state.json"))
if !errors.Is(err, ErrLockHeld) || second != nil {
t.Fatalf("second lock = %#v, %v; want nil, ErrLockHeld", second, err)
}
}
func TestAdvisoryLockCrashReleasesAndReacquires(t *testing.T) {
statePath := filepath.Join(t.TempDir(), "update-state.json")
if os.Getenv("THT_LOCK_CRASH_HELPER") == "1" {
lock, err := acquireLock(os.Getenv("THT_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(), "THT_LOCK_CRASH_HELPER=1", "THT_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")
lockPath := lifecycleLockPath(statePath)
if err := os.WriteFile(lockPath, nil, 0o600); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(lockPath+".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()
}
+967
View File
@@ -0,0 +1,967 @@
package pi
import (
"context"
"crypto/sha256"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"regexp"
"runtime"
"sort"
"strings"
"time"
"github.com/aritmolab/thothii/tools/tht/internal/compose"
"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.-]+)?$`)
)
// Source chooses whether the candidate is built from this checkout or pulled from an immutable image.
type Source string
const (
BuildSource Source = "build"
PullSource Source = "pull"
)
// Request contains only non-secret operator inputs.
type Request struct {
StatePath string
RestartStatePath string
Version string
Source Source
Image string
Confirm bool
Drain bool
}
// Result summarizes the completed, failed, or recovered transaction without command output.
type Result struct {
Phase Phase
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")
}
if request.RestartStatePath == "" {
request.RestartStatePath = pairedRestartStatePath(request.StatePath)
}
if err := validateRestartStatePaths(request.RestartStatePath, request.StatePath); err != nil {
return Result{StatePath: request.StatePath}, err
}
lock, err := acquireLock(request.StatePath)
if err != nil {
return Result{StatePath: request.StatePath}, err
}
defer lock.Release()
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 != 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)
}
request.Image = canonical
}
if err := prepareLifecycleMutation(request.RestartStatePath, hooks.removeFile); err != nil {
return Result{StatePath: request.StatePath}, err
}
if old, err := readState(request.StatePath); err == nil && stateNeedsRecovery(old) {
return Result{StatePath: request.StatePath}, ErrInterruptedUpdate
} else if err != nil && !errors.Is(err, os.ErrNotExist) {
return Result{StatePath: request.StatePath}, err
} else if err == nil && !old.MutationStarted {
if cleanupErr := hooks.removeFile(lifecycleOverridePath(request.StatePath, old.Transaction)); cleanupErr != nil {
return Result{StatePath: request.StatePath}, errors.New("safe prior preparation state could not be cleaned up")
}
}
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}
retErr = errors.Join(retErr, fmt.Errorf("maintenance admission gate could not be cleared: %w", clearErr))
}
}()
if err := waitForInactiveSessions(ctx, runner, request.Drain, hooks.sleep); err != nil {
return Result{StatePath: request.StatePath}, err
}
if err := Doctor(ctx, runner); err != nil {
return Result{StatePath: request.StatePath}, err
}
currentVersion, err := Status(ctx, runner)
if err != nil {
return Result{StatePath: request.StatePath}, err
}
if currentVersion == request.Version {
return Result{Phase: PhaseNoop, StatePath: request.StatePath}, nil
}
configured, err := renderedCore(ctx, runner)
if err != nil {
return Result{StatePath: request.StatePath}, err
}
previous, err := runningImage(ctx, runner, configured.Reference)
if err != nil {
return Result{StatePath: request.StatePath}, err
}
previous.ConfigurationSHA = configured.ConfigurationSHA
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 := hooks.writeState(request.StatePath, state); err != nil {
result, retErr, clearMaintenance = failPreparation(request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
if err := prepareCandidate(ctx, lifecycle, request, candidateReference); err != nil {
result, retErr, clearMaintenance = failPreparation(request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
running, err := activeSessions(ctx, runner)
if err != nil {
result, retErr, clearMaintenance = failPreparation(request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
if running {
result, retErr, clearMaintenance = failPreparation(request.StatePath, overridePath, state, ErrActiveSessions, hooks)
return result, retErr
}
state.MutationStarted = true
if err := hooks.writeState(request.StatePath, state); err != nil {
state.MutationStarted = false
result, retErr, clearMaintenance = failPreparation(request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
clearMaintenance = false
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, lifecycle, candidateReference)
if err != nil {
result, retErr, clearMaintenance = compensate(ctx, runner, request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
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 = PhasePromoting, ""
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 := promoteLifecycleOverride(overridePath, currentImageOverridePath(request.StatePath), candidateReference); err != nil {
result, retErr, clearMaintenance = compensate(ctx, runner, request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
state.Phase = PhaseVerified
if err := hooks.writeState(request.StatePath, state); err != nil {
result, retErr, clearMaintenance = compensate(ctx, runner, request.StatePath, overridePath, state, err, hooks)
return result, retErr
}
clearMaintenance = true
return Result{Phase: PhaseVerified, StatePath: request.StatePath}, nil
}
func stateNeedsRecovery(state State) bool {
switch state.Phase {
case PhaseVerified, PhaseRolledBack, PhaseNoop:
return false
case PhaseFailed:
return state.MutationStarted
default:
return state.MutationStarted
}
}
func failPreparation(statePath, overridePath string, state State, cause error, hooks lifecycleHooks) (Result, error, bool) {
state.Phase = PhaseFailed
state.MutationStarted = false
state.Error = "candidate preparation failed before core mutation"
writeErr := hooks.writeState(statePath, state)
removeErr := hooks.removeFile(overridePath)
message := "candidate preparation failed before core mutation"
if writeErr != nil || removeErr != nil {
message += "; safe preparation cleanup was incomplete"
}
return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("%s: %w", message, cause), true
}
// Rollback restores the image recorded in durable update state. It is safe for interrupted runs.
func Rollback(ctx context.Context, runner Runner, statePath, restartStatePath string, confirm bool) (result Result, retErr error) {
return rollbackWithHooks(ctx, runner, statePath, restartStatePath, confirm, defaultLifecycleHooks)
}
func rollbackWithHooks(ctx context.Context, runner Runner, statePath, restartStatePath string, confirm bool, hooks lifecycleHooks) (result Result, retErr error) {
if restartStatePath == "" {
restartStatePath = pairedRestartStatePath(statePath)
}
if err := validateRestartStatePaths(restartStatePath, statePath); err != nil {
return Result{StatePath: statePath}, err
}
lock, err := acquireLock(statePath)
if err != nil {
return Result{StatePath: statePath}, err
}
defer lock.Release()
if !confirm {
return Result{StatePath: statePath}, ErrConfirmationRequired
}
if err := prepareLifecycleMutation(restartStatePath, hooks.removeFile); err != nil {
return Result{StatePath: statePath}, err
}
maintenanceErr := ensureMaintenance(ctx, runner)
clearMaintenance := maintenanceErr == nil
defer func() {
if !clearMaintenance {
return
}
if clearErr := setMaintenance(context.Background(), runner, false); clearErr != nil {
result = Result{Phase: PhaseFailed, StatePath: statePath}
retErr = errors.Join(retErr, fmt.Errorf("maintenance admission gate could not be cleared: %w", clearErr))
}
}()
if maintenanceErr == nil {
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 maintenanceErr == 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
}
lifecycle := composeOverrideRunner{Runner: runner, path: overridePath}
if maintenanceErr != nil {
stopped, stopErr := coreIsStopped(ctx, runner)
if stopErr != nil || !stopped {
return Result{StatePath: statePath}, maintenanceErr
}
if err := persistMaintenanceWithoutLiveCore(ctx, lifecycle); err != nil {
return Result{Phase: PhaseFailed, StatePath: statePath}, err
}
}
clearMaintenance = false
if err := restore(ctx, lifecycle, state.Previous); err != nil {
state.Phase, state.Error = PhaseFailed, "rollback failed"
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
}
if active, err := activeSessions(ctx, runner); err != nil {
return Result{Phase: PhaseFailed, StatePath: statePath}, err
} else if active {
return Result{Phase: PhaseFailed, StatePath: statePath}, ErrActiveSessions
}
if err := promoteLifecycleOverride(overridePath, currentImageOverridePath(statePath), state.Previous.Reference); err != nil {
return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("rollback restored the core but durable current-image promotion failed: %w", err)
}
state.Phase, state.Error = PhaseRolledBack, ""
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")
}
clearMaintenance = true
return Result{Phase: PhaseRolledBack, StatePath: statePath}, nil
}
func compensate(ctx context.Context, runner Runner, statePath, overridePath string, state State, cause error, hooks lifecycleHooks) (Result, error, bool) {
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 err := ensureMaintenance(context.Background(), runner); err != nil {
stopped, stopErr := coreIsStopped(context.Background(), runner)
if stopErr != nil || !stopped {
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 markerErr := persistMaintenanceWithoutLiveCore(context.Background(), lifecycle); markerErr != nil {
state.Phase, state.Error = PhaseFailed, "maintenance recovery failed"
_ = hooks.writeState(statePath, state)
return Result{Phase: PhaseFailed, StatePath: statePath}, errors.New("update failed and durable maintenance could not be established: recovery required"), false
}
} else 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 restoreErr := restore(ctx, lifecycle, state.Previous); restoreErr != nil {
state.Phase, state.Error = PhaseFailed, "candidate verification and automatic rollback failed"
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"), false
}
if active, err := activeSessions(context.Background(), runner); err != nil || active {
state.Phase, state.Error = PhaseFailed, "restored rollback inventory failed"
_ = hooks.writeState(statePath, state)
return Result{Phase: PhaseFailed, StatePath: statePath}, errors.New("previous core image was restored but rollback inventory is not quiescent: recovery required"), false
}
if err := promoteLifecycleOverride(overridePath, currentImageOverridePath(statePath), state.Previous.Reference); err != nil {
return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("previous core image was restored but durable selector promotion failed: %w", err), false
}
state.Phase, state.Error = PhaseRolledBack, ""
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: %w", cause), true
}
func coreIsStopped(ctx context.Context, runner Runner) (bool, error) {
result, err := runCompose(ctx, runner, "ps", "--status", "running", "-q", "core")
if err != nil {
return false, commandError("core running-state check", result, err)
}
return strings.TrimSpace(result.Stdout) == "", nil
}
const maintenanceMarkerScript = `
const fs = require("node:fs");
const path = require("node:path");
const marker = process.env.THT_MAINTENANCE_FILE;
if (!marker) throw new Error("THT_MAINTENANCE_FILE is required");
const directory = path.dirname(marker);
fs.mkdirSync(directory, { recursive: true });
const temporary = marker + ".rollback-" + process.pid + "-" + Date.now();
let file;
try {
file = fs.openSync(temporary, "wx", 0o600);
fs.writeFileSync(file, "{\"version\":1,\"active\":true}\n", "utf8");
fs.fsyncSync(file);
fs.closeSync(file);
file = undefined;
fs.renameSync(temporary, marker);
const directoryFile = fs.openSync(directory, "r");
try { fs.fsyncSync(directoryFile); } finally { fs.closeSync(directoryFile); }
} catch (error) {
if (file !== undefined) try { fs.closeSync(file); } catch {}
try { fs.unlinkSync(temporary); } catch {}
throw error;
}
`
func persistMaintenanceWithoutLiveCore(ctx context.Context, runner Runner) error {
result, err := runCompose(ctx, runner,
"run", "--rm", "--no-deps", "--entrypoint", "node", "core", "-e", maintenanceMarkerScript,
)
if err != nil {
return commandError("durable maintenance recovery", result, err)
}
return nil
}
func sourceValue(request Request) string {
if request.Source == PullSource {
return request.Image
}
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
}
func setMaintenance(ctx context.Context, runner Runner, enabled bool) error {
path := "deactivate"
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...)
status, valid := parseMaintenanceStatus(result.Stdout)
if err == nil && valid && status.Active == enabled && status.Admissions == 0 && !status.RecoveryRequired {
return nil
}
// Status identifies the safest immediate state after an ambiguous response. It cannot
// acknowledge durability for an operation whose command returned an error.
observed, statusErr := MaintenanceStatus(ctx, runner)
if valid && status.RecoveryRequired {
return recoveryRequired("maintenance durability was explicitly not acknowledged", err)
}
if statusErr == nil && observed.RecoveryRequired {
return recoveryRequired("maintenance durability was explicitly not acknowledged", err)
}
if err != nil {
if statusErr == nil && observed.Active == enabled && observed.Admissions == 0 {
return recoveryRequired("maintenance durability was not acknowledged after a failed command", err)
}
return commandError("maintenance admission gate", result, err)
}
if statusErr == nil && observed.Active == enabled && observed.Admissions == 0 {
return nil
}
return errors.New("maintenance admission gate did not acknowledge a quiescent state")
}
type MaintenanceState struct {
Active bool `json:"active"`
Admissions int `json:"admissions"`
RecoveryRequired bool `json:"recoveryRequired"`
}
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 && !status.RecoveryRequired {
return nil
}
if err == nil && status.RecoveryRequired {
return recoveryRequired("maintenance durability was explicitly not acknowledged", nil)
}
return setMaintenance(ctx, runner, true)
}
func activeSessions(ctx context.Context, runner Runner) (bool, error) {
scope := "all"
if scoped, ok := runner.(interface{ SessionInventoryScope() string }); ok {
if requested := scoped.SessionInventoryScope(); requested == "mine" || requested == "all" {
scope = requested
}
}
args := append([]string{"exec", "-T", "core", "curl", "-fsS"}, internalIdentityHeaders...)
args = append(args, "http://127.0.0.1:8787/sessions?scope="+scope)
result, err := runCompose(ctx, runner, args...)
if err != nil {
return false, commandError("active-session check", result, err)
}
var payload []struct {
Status string `json:"status"`
Archived bool `json:"archived"`
}
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 {
if !session.Archived && session.Status != "finalized" && session.Status != "closed" {
return true, nil
}
}
return false, nil
}
func waitForInactiveSessions(ctx context.Context, runner Runner, drain bool, sleep func(time.Duration)) error {
running, err := activeSessions(ctx, runner)
if err != nil {
return err
}
if !running {
return nil
}
if !drain {
return ErrActiveSessions
}
for attempts := 0; attempts < 30; attempts++ {
running, err = activeSessions(ctx, runner)
if err != nil {
return err
}
if !running {
return nil
}
sleep(time.Second)
}
return ErrActiveSessions
}
func runningImage(ctx context.Context, runner Runner, reference string) (Image, error) {
container, err := runCompose(ctx, runner, "ps", "-q", "core")
if err != nil || strings.TrimSpace(container.Stdout) == "" {
return Image{}, commandError("running core image check", container, err)
}
id := strings.TrimSpace(container.Stdout)
image, err := runner.Run(ctx, []string{"inspect", "--format", "{{.Image}}", id}, nil)
if err != nil || strings.TrimSpace(image.Stdout) == "" {
return Image{}, commandError("running core image check", image, err)
}
mounts, err := runner.Run(ctx, []string{"inspect", "--format", "{{json .Mounts}}", id}, nil)
if err != nil {
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"`
}
if err := json.Unmarshal([]byte(mounts.Stdout), &raw); err != nil {
return Image{}, errors.New("core returned invalid persistence mount data")
}
if len(raw) == 0 {
return Image{}, errors.New("core has no persistence mounts to preserve")
}
contract := make([]Mount, 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), SourceAliases: mountSourceAliases(mount.Type, mount.Source, runtime.GOOS), 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, Mounts: contract, MountFingerprint: mountFingerprint(contract)}, nil
}
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 {
return commandError("Pi image build", result, err)
}
return nil
}
pull, err := runner.Run(ctx, []string{"pull", request.Image}, nil)
if err != nil {
return commandError("Pi image pull", pull, err)
}
return tagImage(ctx, runner, request.Image, candidateReference, "Pi image tag")
}
func recreateCore(ctx context.Context, runner Runner) error {
result, err := runCompose(ctx, runner, "up", "--detach", "--wait", "--wait-timeout", "45", "--no-deps", "--force-recreate", "core")
if err != nil {
return commandError("core recreation", result, err)
}
return nil
}
func verifyCandidate(ctx context.Context, runner Runner, wanted string, previous Image) error {
health, err := runCompose(ctx, runner, "exec", "-T", "core", "curl", "-fsS", "http://127.0.0.1:8787/health")
if err != nil {
return commandError("core health check", health, err)
}
if err := verifyCandidateVersionIdentity(ctx, runner, wanted); err != nil {
return err
}
if err := Test(ctx, runner); err != nil {
return err
}
configured, err := renderedCore(ctx, runner)
if err != nil {
return err
}
if configured.ConfigurationSHA != previous.ConfigurationSHA {
return errors.New("external endpoint configuration changed during Pi update")
}
after, err := runningImage(ctx, runner, configured.Reference)
if err != nil {
return err
}
if !sameMounts(previous.Mounts, after.Mounts) {
return errors.New("core persistence mount contract changed during Pi update")
}
return nil
}
func verifyCandidateVersionIdentity(ctx context.Context, runner Runner, wanted string) error {
executable, err := Status(ctx, runner)
if err != nil {
return err
}
environment, label, err := expectedVersions(ctx, runner)
if err != nil {
return err
}
if executable != wanted || environment != wanted || label != wanted {
return errors.New("candidate Pi executable, PI_VERSION, and image label do not all match the requested pinned version")
}
return nil
}
func restore(ctx context.Context, runner Runner, previous Image) error {
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
}
after, err := runningImage(ctx, runner, configured.Reference)
if err != nil {
return err
}
if after.ID != previous.ID {
return errors.New("rollback core image does not match recorded previous image")
}
if configured.ConfigurationSHA != previous.ConfigurationSHA {
return errors.New("external endpoint configuration drift prevents rollback proof")
}
if !sameMounts(previous.Mounts, after.Mounts) {
return errors.New("core persistence mount contract changed during rollback")
}
if err := Doctor(ctx, runner); err != nil {
return err
}
if err := Test(ctx, runner); err != nil {
return err
}
return nil
}
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)
}
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:tht-" + transaction + "-" + role
}
func lifecycleOverridePath(statePath, transaction string) string {
if transaction == "" {
transaction = "recovery"
}
return filepath.Join(filepath.Dir(statePath), "pi-lifecycle-"+transaction+".yaml")
}
func currentImageOverridePath(statePath string) string {
return filepath.Join(filepath.Dir(statePath), "current-image.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:
core:
image: ` + string(quoted) + `
workspace-maintenance:
image: ` + string(quoted) + `
`)
if err := writeFileDurably(path, ".pi-lifecycle-", contents); err != nil {
return errors.New("lifecycle image override could not be written durably")
}
return nil
}
func promoteLifecycleOverride(source, destination, expectedImage string) error {
return promoteLifecycleOverrideWith(source, destination, expectedImage, durableReplace)
}
func promoteLifecycleOverrideWith(
source, destination, expectedImage string,
replace func(string, string, string) error,
) error {
if err := replace(source, destination, filepath.Dir(destination)); err != nil {
selected, readErr := readLifecycleOverride(destination)
if readErr == nil && selected == expectedImage {
return recoveryRequired("lifecycle image override changed but durability was not acknowledged", err)
}
return recoveryRequired("lifecycle image override could not be promoted durably", err)
}
selected, err := readLifecycleOverride(destination)
if err != nil || selected != expectedImage {
return errors.New("promoted lifecycle image override could not be verified")
}
return nil
}
func readLifecycleOverride(path string) (string, error) {
contents, err := os.ReadFile(path)
if err != nil {
return "", err
}
for _, line := range strings.Split(string(contents), "\n") {
line = strings.TrimSpace(line)
if !strings.HasPrefix(line, "image:") {
continue
}
encoded := strings.TrimSpace(strings.TrimPrefix(line, "image:"))
var image string
if json.Unmarshal([]byte(encoded), &image) != nil || image == "" || strings.ContainsAny(image, "\r\n") {
return "", errors.New("lifecycle image override is invalid")
}
return image, nil
}
return "", errors.New("lifecycle image override has no core image")
}
// 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()
return recoverMaintenanceLocked(ctx, runner, statePath)
}
func recoverMaintenanceLocked(ctx context.Context, runner Runner, statePath string) error {
state, stateErr := readState(statePath)
if stateErr == nil {
transactionOverride := lifecycleOverridePath(statePath, state.Transaction)
switch {
case state.Phase == PhasePromoting:
if err := recoverPromotion(ctx, runner, statePath, transactionOverride, &state); err != nil {
return err
}
case !state.MutationStarted && state.Phase != PhaseVerified && state.Phase != PhaseRolledBack && state.Phase != PhaseNoop:
state.Phase, state.Error = PhaseFailed, "candidate preparation interrupted before core mutation"
if err := writeState(statePath, state); err != nil {
return errors.New("maintenance recovery could not finalize safe preparation state")
}
if err := durableRemove(transactionOverride); err != nil {
return errors.New("maintenance recovery could not remove the safe preparation override")
}
case stateNeedsRecovery(state):
return ErrInterruptedUpdate
default:
if err := durableRemove(transactionOverride); 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 recoverPromotion(ctx context.Context, runner Runner, statePath, transactionOverride string, state *State) error {
currentOverride := currentImageOverridePath(statePath)
selected, currentErr := readLifecycleOverride(currentOverride)
if currentErr != nil || selected != state.Candidate.Reference {
pending, pendingErr := readLifecycleOverride(transactionOverride)
if pendingErr != nil || pending != state.Candidate.Reference {
if currentErr == nil && selected == state.Previous.Reference {
if err := verifyRestoredCurrent(ctx, runner, state.Previous); err != nil {
return ErrInterruptedUpdate
}
state.Phase, state.Error = PhaseRolledBack, ""
return writeState(statePath, *state)
}
return ErrInterruptedUpdate
}
if err := promoteLifecycleOverride(transactionOverride, currentOverride, state.Candidate.Reference); err != nil {
return err
}
}
if err := verifyCandidate(ctx, runner, state.Target.Version, state.Previous); err != nil {
return err
}
state.Phase, state.Error = PhaseVerified, ""
return writeState(statePath, *state)
}
func verifyRestoredCurrent(ctx context.Context, runner Runner, previous Image) error {
configured, err := renderedCore(ctx, runner)
if err != nil {
return err
}
after, err := runningImage(ctx, runner, configured.Reference)
if err != nil {
return err
}
if after.ID != previous.ID || configured.ConfigurationSHA != previous.ConfigurationSHA || !sameMounts(previous.Mounts, after.Mounts) {
return errors.New("running core does not match the durable previous-image selector")
}
if err := Doctor(ctx, runner); err != nil {
return err
}
return Test(ctx, runner)
}
func sameStrings(left, right []string) bool {
left, right = append([]string(nil), left...), append([]string(nil), right...)
sort.Strings(left)
sort.Strings(right)
return strings.Join(left, "\x00") == strings.Join(right, "\x00")
}
func sameMounts(left, right []Mount) bool {
if len(left) != len(right) {
return false
}
identityWithoutSource := func(m Mount) string {
return m.Type + "\x00" + m.Name + "\x00" + m.Destination + "\x00" + fmt.Sprint(m.RW) + "\x00" + m.Options
}
sourceMatches := func(a, b Mount) bool {
if a.SourceSHA256 == b.SourceSHA256 {
return true
}
for _, alias := range a.SourceAliases {
if alias == b.SourceSHA256 {
return true
}
}
for _, alias := range b.SourceAliases {
if alias == a.SourceSHA256 {
return true
}
}
return false
}
matched := make([]bool, len(right))
for _, candidate := range left {
found := false
for index, observed := range right {
if matched[index] || identityWithoutSource(candidate) != identityWithoutSource(observed) || !sourceMatches(candidate, observed) {
continue
}
matched[index], found = true, true
break
}
if !found {
return false
}
}
return true
}
File diff suppressed because it is too large Load Diff
+102
View File
@@ -0,0 +1,102 @@
// Package safeio reads and writes local files without following symlinked path components.
package safeio
import (
"errors"
"io"
"os"
"path/filepath"
"strings"
"unicode/utf8"
)
var ErrUnsafeFile = errors.New("unsafe file")
// ValidateCanonicalPath rejects relative or lexically non-canonical paths before they are opened.
func ValidateCanonicalPath(path string) error {
if !filepath.IsAbs(path) || filepath.Clean(path) != path || strings.Contains(path, string(filepath.Separator)+".."+string(filepath.Separator)) {
return ErrUnsafeFile
}
return nil
}
func readBoundedRegularFile(path string, file *os.File, maximum int64) ([]byte, error) {
if maximum < 0 || maximum == int64(^uint64(0)>>1) {
return nil, ErrUnsafeFile
}
info, err := file.Stat()
if err != nil || !info.Mode().IsRegular() || !hasSingleLink(info) {
return nil, ErrUnsafeFile
}
contents, err := io.ReadAll(io.LimitReader(file, maximum+1))
if err != nil || int64(len(contents)) > maximum {
return nil, ErrUnsafeFile
}
after, err := file.Stat()
if err != nil || !after.Mode().IsRegular() || !hasSingleLink(after) || !os.SameFile(info, after) {
return nil, ErrUnsafeFile
}
current, err := os.Stat(path)
if err != nil || !os.SameFile(info, current) {
return nil, ErrUnsafeFile
}
return contents, nil
}
func ReadCanonicalUTF8(path string, maximum int64) (string, error) {
contents, err := ReadCanonicalRegular(path, maximum)
if err != nil {
return "", err
}
if !utf8.Valid(contents) {
return "", ErrUnsafeFile
}
return string(contents), nil
}
func WriteCanonicalNewFile(path string, contents []byte, mode os.FileMode) error {
if err := ValidateCanonicalPath(path); err != nil {
return err
}
parent := filepath.Dir(path)
if err := requireCanonicalDirectory(parent); err != nil {
return err
}
if info, err := os.Lstat(path); err == nil {
if !info.Mode().IsRegular() || info.Mode()&os.ModeSymlink != 0 || info.Mode()&os.ModeType != 0 {
return ErrUnsafeFile
}
return ErrUnsafeFile
} else if !errors.Is(err, os.ErrNotExist) {
return ErrUnsafeFile
}
file, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_EXCL, mode)
if err != nil {
return ErrUnsafeFile
}
defer file.Close()
if _, err := file.Write(contents); err != nil {
_ = os.Remove(path)
return ErrUnsafeFile
}
if err := file.Sync(); err != nil {
_ = os.Remove(path)
return ErrUnsafeFile
}
return nil
}
func requireCanonicalDirectory(path string) error {
if err := ValidateCanonicalPath(path); err != nil {
return err
}
resolved, err := filepath.EvalSymlinks(path)
if err != nil || resolved != path {
return ErrUnsafeFile
}
info, err := os.Stat(path)
if err != nil || !info.IsDir() {
return ErrUnsafeFile
}
return nil
}
+95
View File
@@ -0,0 +1,95 @@
package safeio
import (
"errors"
"os"
"path/filepath"
"testing"
"github.com/aritmolab/thothii/tools/tht/internal/testsupport"
)
func TestReadCanonicalRegularRejectsFinalAndParentSymlinks(t *testing.T) {
temporaryRoot, err := filepath.EvalSymlinks(os.TempDir())
if err != nil {
t.Fatal(err)
}
root, err := os.MkdirTemp(temporaryRoot, "tht-safeio-")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.RemoveAll(root) })
realDirectory := filepath.Join(root, "real")
if err := os.Mkdir(realDirectory, 0o700); err != nil {
t.Fatal(err)
}
realFile := filepath.Join(realDirectory, "secret")
if err := os.WriteFile(realFile, []byte("secret"), 0o600); err != nil {
t.Fatal(err)
}
parentLink := filepath.Join(root, "parent-link")
testsupport.SymlinkOrSkip(t, realDirectory, parentLink)
if _, err := ReadCanonicalRegular(filepath.Join(parentLink, "secret"), 1024); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("parent symlink error = %v, want ErrUnsafeFile", err)
}
finalLink := filepath.Join(root, "final-link")
testsupport.SymlinkOrSkip(t, realFile, finalLink)
if _, err := ReadCanonicalRegular(finalLink, 1024); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("final symlink error = %v, want ErrUnsafeFile", err)
}
}
func TestReadCanonicalUTF8RejectsNonUTF8AndHardlinks(t *testing.T) {
temporaryRoot, err := filepath.EvalSymlinks(os.TempDir())
if err != nil {
t.Fatal(err)
}
root, err := os.MkdirTemp(temporaryRoot, "tht-safeio-")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.RemoveAll(root) })
nonUTF8 := filepath.Join(root, "annotations.yaml")
if err := os.WriteFile(nonUTF8, []byte{0xff, 0xfe, 0xfd}, 0o600); err != nil {
t.Fatal(err)
}
if _, err := ReadCanonicalUTF8(nonUTF8, 1024); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("ReadCanonicalUTF8(nonUTF8) error = %v, want ErrUnsafeFile", err)
}
target := filepath.Join(root, "regular.txt")
if err := os.WriteFile(target, []byte("linked"), 0o600); err != nil {
t.Fatal(err)
}
link := filepath.Join(root, "hardlink.txt")
if err := os.Link(target, link); err != nil {
t.Fatal(err)
}
if _, err := ReadCanonicalUTF8(link, 1024); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("ReadCanonicalUTF8(hardlink) error = %v, want ErrUnsafeFile", err)
}
}
func TestWriteCanonicalNewFileRejectsExistingTargets(t *testing.T) {
temporaryRoot, err := filepath.EvalSymlinks(os.TempDir())
if err != nil {
t.Fatal(err)
}
root, err := os.MkdirTemp(temporaryRoot, "tht-safeio-")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.RemoveAll(root) })
path := filepath.Join(root, "artifact.yaml")
if err := os.WriteFile(path, []byte("existing"), 0o600); err != nil {
t.Fatal(err)
}
if err := WriteCanonicalNewFile(path, []byte("new"), 0o600); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("WriteCanonicalNewFile(existing) error = %v, want ErrUnsafeFile", err)
}
}
+57
View File
@@ -0,0 +1,57 @@
//go:build !windows
package safeio
import (
"os"
"strings"
"golang.org/x/sys/unix"
)
// ReadCanonicalRegular opens an absolute canonical path component by component from the root
// descriptor. O_NOFOLLOW rejects symlinks at every component, and the open directory descriptors
// prevent later parent replacement from redirecting the final open.
func ReadCanonicalRegular(path string, maximum int64) ([]byte, error) {
if err := ValidateCanonicalPath(path); err != nil {
return nil, err
}
components := strings.Split(strings.TrimPrefix(path, string(os.PathSeparator)), string(os.PathSeparator))
if len(components) == 0 || components[0] == "" {
return nil, ErrUnsafeFile
}
directory, err := unix.Open(string(os.PathSeparator), unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY, 0)
if err != nil {
return nil, ErrUnsafeFile
}
directories := []int{directory}
defer func() { closeUnixDescriptors(directories) }()
for _, component := range components[:len(components)-1] {
nextDirectory, err := unix.Openat(directory, component, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY|unix.O_NOFOLLOW, 0)
if err != nil {
return nil, ErrUnsafeFile
}
directory = nextDirectory
directories = append(directories, directory)
}
descriptor, err := unix.Openat(directory, components[len(components)-1], unix.O_RDONLY|unix.O_CLOEXEC|unix.O_NOFOLLOW|unix.O_NONBLOCK, 0)
if err != nil {
return nil, ErrUnsafeFile
}
file := os.NewFile(uintptr(descriptor), "tht-safeio")
if file == nil {
unix.Close(descriptor)
return nil, ErrUnsafeFile
}
defer file.Close()
return readBoundedRegularFile(path, file, maximum)
}
func closeUnixDescriptors(descriptors []int) {
for _, descriptor := range descriptors {
unix.Close(descriptor)
}
}
@@ -0,0 +1,32 @@
//go:build !windows
package safeio
import (
"errors"
"os"
"path/filepath"
"testing"
"golang.org/x/sys/unix"
)
func TestReadCanonicalRegularRejectsNamedPipeWithoutBlocking(t *testing.T) {
temporaryRoot, err := filepath.EvalSymlinks(os.TempDir())
if err != nil {
t.Fatal(err)
}
root, err := os.MkdirTemp(temporaryRoot, "tht-safeio-")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.RemoveAll(root) })
pipe := filepath.Join(root, "secret-pipe")
if err := unix.Mkfifo(pipe, 0o600); err != nil {
t.Fatal(err)
}
if _, err := ReadCanonicalRegular(pipe, 1024); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("named pipe error = %v, want ErrUnsafeFile", err)
}
}
@@ -0,0 +1,96 @@
//go:build windows
package safeio
import (
"os"
"path/filepath"
"strings"
"golang.org/x/sys/windows"
)
const windowsRetainedHandleShareMode uint32 = windows.FILE_SHARE_READ | windows.FILE_SHARE_WRITE
// ReadCanonicalRegular opens each component with FILE_FLAG_OPEN_REPARSE_POINT and rejects a
// reparse point on the opened handle before opening the next component. Retained handles allow
// ordinary read/write sharing but deny delete sharing, which blocks rename or deletion after a
// component is opened and throughout the final read. Windows' Win32 API does not expose a
// portable descriptor-relative equivalent of POSIX openat, so a hostile local actor can still
// replace a not-yet-opened normal component between absolute-path opens. Installation directories
// therefore need trusted local filesystem/ACL ownership on Windows.
func ReadCanonicalRegular(path string, maximum int64) ([]byte, error) {
if err := ValidateCanonicalPath(path); err != nil {
return nil, err
}
volume := filepath.VolumeName(path)
root := volume + string(filepath.Separator)
components := strings.Split(strings.TrimPrefix(path, root), string(filepath.Separator))
if volume == "" || len(components) == 0 || components[0] == "" {
return nil, ErrUnsafeFile
}
current := root
parents := make([]windows.Handle, 0, len(components)-1)
defer func() { closeWindowsHandles(parents) }()
for _, component := range components[:len(components)-1] {
current = filepath.Join(current, component)
handle, err := openWindowsComponent(current, true)
if err != nil {
return nil, ErrUnsafeFile
}
parents = append(parents, handle)
}
current = filepath.Join(current, components[len(components)-1])
handle, err := openWindowsComponent(current, false)
if err != nil {
return nil, ErrUnsafeFile
}
file := os.NewFile(uintptr(handle), "tht-safeio")
if file == nil {
windows.CloseHandle(handle)
return nil, ErrUnsafeFile
}
defer file.Close()
return readBoundedRegularFile(path, file, maximum)
}
func openWindowsComponent(path string, directory bool) (windows.Handle, error) {
flags := uint32(windows.FILE_FLAG_OPEN_REPARSE_POINT)
if directory {
flags |= windows.FILE_FLAG_BACKUP_SEMANTICS
} else {
flags |= windows.FILE_ATTRIBUTE_NORMAL
}
handle, err := windows.CreateFile(
windows.StringToUTF16Ptr(path),
windows.GENERIC_READ,
windowsRetainedHandleShareMode,
nil,
windows.OPEN_EXISTING,
flags,
0,
)
if err != nil {
return 0, err
}
var information windows.ByHandleFileInformation
if err := windows.GetFileInformationByHandle(handle, &information); err != nil {
windows.CloseHandle(handle)
return 0, err
}
if information.FileAttributes&windows.FILE_ATTRIBUTE_REPARSE_POINT != 0 ||
(directory && information.FileAttributes&windows.FILE_ATTRIBUTE_DIRECTORY == 0) ||
(!directory && information.FileAttributes&windows.FILE_ATTRIBUTE_DIRECTORY != 0) {
windows.CloseHandle(handle)
return 0, ErrUnsafeFile
}
return handle, nil
}
func closeWindowsHandles(handles []windows.Handle) {
for _, handle := range handles {
windows.CloseHandle(handle)
}
}
@@ -0,0 +1,68 @@
//go:build windows
package safeio
import (
"os"
"path/filepath"
"testing"
"golang.org/x/sys/windows"
)
const expectedWindowsRetainedHandleShareMode = windows.FILE_SHARE_READ | windows.FILE_SHARE_WRITE
// Keep this contract compile-enforced so Windows cross-test compilation catches a future
// FILE_SHARE_DELETE regression even when the tests are compiled on a non-Windows host.
var _ [windowsRetainedHandleShareMode - expectedWindowsRetainedHandleShareMode]struct{}
var _ [expectedWindowsRetainedHandleShareMode - windowsRetainedHandleShareMode]struct{}
func TestOpenWindowsComponentBlocksMutationWhileHandleIsRetained(t *testing.T) {
t.Run("parent rename", func(t *testing.T) {
parent := filepath.Join(t.TempDir(), "parent")
if err := os.Mkdir(parent, 0o700); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(parent, "secret"), []byte("secret"), 0o600); err != nil {
t.Fatal(err)
}
handle, err := openWindowsComponent(parent, true)
if err != nil {
t.Fatal(err)
}
renamed := parent + "-renamed"
if err := os.Rename(parent, renamed); err == nil {
windows.CloseHandle(handle)
t.Fatal("parent rename succeeded while its safe-I/O handle was retained")
}
if err := windows.CloseHandle(handle); err != nil {
t.Fatal(err)
}
if err := os.Rename(parent, renamed); err != nil {
t.Fatalf("parent rename after closing its safe-I/O handle: %v", err)
}
})
t.Run("final delete", func(t *testing.T) {
path := filepath.Join(t.TempDir(), "secret")
if err := os.WriteFile(path, []byte("secret"), 0o600); err != nil {
t.Fatal(err)
}
handle, err := openWindowsComponent(path, false)
if err != nil {
t.Fatal(err)
}
if err := os.Remove(path); err == nil {
windows.CloseHandle(handle)
t.Fatal("final-file deletion succeeded while its safe-I/O handle was retained")
}
if err := windows.CloseHandle(handle); err != nil {
t.Fatal(err)
}
if err := os.Remove(path); err != nil {
t.Fatalf("final-file deletion after closing its safe-I/O handle: %v", err)
}
})
}
@@ -0,0 +1,13 @@
//go:build !windows
package safeio
import (
"os"
"syscall"
)
func hasSingleLink(info os.FileInfo) bool {
stat, ok := info.Sys().(*syscall.Stat_t)
return ok && stat.Nlink == 1
}
@@ -0,0 +1,9 @@
//go:build windows
package safeio
import "os"
func hasSingleLink(info os.FileInfo) bool {
return true
}
+390
View File
@@ -0,0 +1,390 @@
// Package serverops implements bounded, installation-aware server maintenance operations.
package serverops
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strconv"
"strings"
"github.com/aritmolab/thothii/tools/tht/internal/compose"
"github.com/aritmolab/thothii/tools/tht/internal/config"
)
var (
ErrConfirmationRequired = errors.New("explicit confirmation is required")
ErrUnsafeState = errors.New("server operation refused in the current state")
)
type Runner interface {
Run(context.Context, []string, io.Reader) (compose.Result, error)
}
type Stage string
const (
StageContainerInspection Stage = "container-inspection"
StageMigrationConfig Stage = "migration-config"
StageMigrationVerification Stage = "migration-config-verification"
StageSessionMigration Stage = "session-migration"
StageContainerRemoval Stage = "container-removal"
StageRemovalVerification Stage = "removal-verification"
)
type ExitClass string
const (
ExitClassNonzero ExitClass = "nonzero-exit"
ExitClassUnavailable ExitClass = "unavailable"
ExitClassTimeout ExitClass = "timeout"
ExitClassInvocation ExitClass = "invocation-failure"
)
// OperationError reports only allowlisted operation metadata from Error. Complete subprocess
// detail is exposed separately so the CLI can redact it before applying its display bound.
type OperationError struct {
stage Stage
class ExitClass
detail string
}
func (e *OperationError) Error() string {
return fmt.Sprintf("stage=%s class=%s", e.stage, e.class)
}
func (e *OperationError) Stage() Stage {
return e.stage
}
func (e *OperationError) Class() ExitClass {
return e.class
}
func (e *OperationError) Detail() string {
return e.detail
}
type MigrationStatus struct {
Applied []string `json:"applied"`
Drifted []string `json:"drifted"`
Pending []string `json:"pending"`
}
type Container struct {
ID string `json:"ID"`
Name string `json:"Name"`
Service string `json:"Service"`
State string `json:"State"`
}
type RemovalResult struct {
Targets []Container
Preserved int
}
// MigrateSessions runs only the one-shot migration service and proves the resulting schema state.
func MigrateSessions(ctx context.Context, installation config.Installation, runner Runner, confirmed bool) (MigrationStatus, error) {
if !confirmed {
return MigrationStatus{}, ErrConfirmationRequired
}
if installation.Profile != "server" {
return MigrationStatus{}, fmt.Errorf("%w: session migration requires a server installation", ErrUnsafeState)
}
containers, err := inspectContainers(ctx, installation, runner, StageContainerInspection)
if err != nil {
return MigrationStatus{}, err
}
if err := requireStopped(containers); err != nil {
return MigrationStatus{}, err
}
rendered, err := runCompose(ctx, runner, StageMigrationConfig, installation.ComposeArgs("--profile", "session-migrate", "config", "--format", "json"))
if err != nil {
return MigrationStatus{}, err
}
coreImage, err := selectedCoreImage(rendered.Stdout)
if err != nil {
return MigrationStatus{}, err
}
override, cleanup, err := migrationOverride(installation, coreImage)
if err != nil {
return MigrationStatus{}, err
}
defer cleanup()
configArgs, err := installation.ComposeArgsWithFinalOverride(override, "--profile", "session-migrate", "config", "--format", "json")
if err != nil {
return MigrationStatus{}, err
}
finalConfig, err := runCompose(ctx, runner, StageMigrationVerification, configArgs)
if err != nil {
return MigrationStatus{}, err
}
if err := requireMigrationImage(finalConfig.Stdout, coreImage); err != nil {
return MigrationStatus{}, err
}
runArgs, err := installation.ComposeArgsWithFinalOverride(
override, "--profile", "session-migrate", "run", "--rm", "--no-deps", "--no-TTY", "session-migrate",
)
if err != nil {
return MigrationStatus{}, err
}
result, err := runCompose(ctx, runner, StageSessionMigration, runArgs)
if err != nil {
return MigrationStatus{}, err
}
status, err := parseMigrationStatus(result.Stdout)
if err != nil {
return MigrationStatus{}, err
}
if len(status.Pending) != 0 || len(status.Drifted) != 0 {
return status, fmt.Errorf("%w: session migration did not finish cleanly", ErrUnsafeState)
}
return status, nil
}
// Remove deletes only the exact stopped core/frontend container IDs displayed by the command.
// A nil confirmation performs inspection only; a non-nil confirmation must equal every target ID.
func Remove(ctx context.Context, installation config.Installation, runner Runner, confirmedIDs []string) (RemovalResult, error) {
if installation.Profile != "server" {
return RemovalResult{}, fmt.Errorf("%w: removal requires a server installation", ErrUnsafeState)
}
targets, err := inspectContainers(ctx, installation, runner, StageContainerInspection)
result := RemovalResult{Targets: targets}
if err != nil {
return result, err
}
if err := requireStopped(targets); err != nil {
return result, err
}
if confirmedIDs == nil {
return result, ErrConfirmationRequired
}
if !sameTargetIDs(targets, confirmedIDs) {
return result, fmt.Errorf("%w: confirmed container IDs differ from current targets", ErrUnsafeState)
}
paths, err := installation.PreservationPaths()
if err != nil {
return result, fmt.Errorf("%w: preservation paths could not be verified", ErrUnsafeState)
}
snapshots, err := snapshotPaths(paths)
if err != nil {
return result, err
}
if len(targets) > 0 {
args := []string{"rm"}
for _, target := range targets {
args = append(args, target.ID)
}
if _, err := runDocker(ctx, runner, StageContainerRemoval, args); err != nil {
return result, err
}
}
remaining, err := inspectContainers(ctx, installation, runner, StageRemovalVerification)
if err != nil {
return result, err
}
if len(remaining) != 0 {
return result, fmt.Errorf("%w: installation containers changed during removal", ErrUnsafeState)
}
if err := verifySnapshots(snapshots); err != nil {
return result, err
}
result.Preserved = len(snapshots)
return result, nil
}
func sameTargetIDs(targets []Container, confirmed []string) bool {
if len(targets) != len(confirmed) {
return false
}
wanted := make(map[string]struct{}, len(confirmed))
for _, id := range confirmed {
if strings.TrimSpace(id) == "" {
return false
}
if _, duplicate := wanted[id]; duplicate {
return false
}
wanted[id] = struct{}{}
}
for _, target := range targets {
if _, exists := wanted[target.ID]; !exists {
return false
}
}
return true
}
func inspectContainers(ctx context.Context, installation config.Installation, runner Runner, stage Stage) ([]Container, error) {
result, err := runCompose(ctx, runner, stage, installation.ComposeArgs("ps", "--all", "--format", "json", "core", "frontend"))
if err != nil {
return nil, err
}
var containers []Container
if err := json.Unmarshal([]byte(result.Stdout), &containers); err != nil {
return nil, fmt.Errorf("%w: Compose returned invalid container status", ErrUnsafeState)
}
seen := make(map[string]struct{})
for _, container := range containers {
if (container.Service != "core" && container.Service != "frontend") || container.ID == "" || container.Name == "" {
return nil, fmt.Errorf("%w: Compose returned an unexpected removal target", ErrUnsafeState)
}
if _, exists := seen[container.ID]; exists {
return nil, fmt.Errorf("%w: Compose returned duplicate container IDs", ErrUnsafeState)
}
seen[container.ID] = struct{}{}
}
return containers, nil
}
func requireStopped(containers []Container) error {
for _, container := range containers {
if strings.ToLower(container.State) != "exited" {
return fmt.Errorf("%w: %s is not stopped", ErrUnsafeState, container.Service)
}
}
return nil
}
func selectedCoreImage(document string) (string, error) {
services, err := renderedServices(document)
if err != nil {
return "", err
}
core, exists := services["core"]
if !exists || strings.TrimSpace(core.Image) == "" {
return "", fmt.Errorf("%w: rendered core image is missing", ErrUnsafeState)
}
if _, exists := services["session-migrate"]; !exists {
return "", fmt.Errorf("%w: rendered migration service is missing", ErrUnsafeState)
}
return core.Image, nil
}
type renderedService struct {
Image string `json:"image"`
Build json.RawMessage `json:"build"`
}
func renderedServices(document string) (map[string]renderedService, error) {
var configDocument struct {
Services map[string]renderedService `json:"services"`
}
if err := json.Unmarshal([]byte(document), &configDocument); err != nil {
return nil, fmt.Errorf("%w: Compose returned invalid rendered configuration", ErrUnsafeState)
}
return configDocument.Services, nil
}
func requireMigrationImage(document, coreImage string) error {
services, err := renderedServices(document)
if err != nil {
return err
}
migrator, exists := services["session-migrate"]
if !exists || migrator.Image != coreImage {
return fmt.Errorf("%w: migration image differs from selected core image", ErrUnsafeState)
}
if len(migrator.Build) != 0 && strings.TrimSpace(string(migrator.Build)) != "null" {
return fmt.Errorf("%w: migration service unexpectedly declares a build", ErrUnsafeState)
}
return nil
}
func migrationOverride(installation config.Installation, image string) (string, func(), error) {
control := installation.ControlDirectory()
if err := os.MkdirAll(control, 0o700); err != nil {
return "", func() {}, errors.New("migration control directory could not be created")
}
info, err := os.Lstat(control)
if err != nil || !info.IsDir() || info.Mode()&os.ModeSymlink != 0 {
return "", func() {}, errors.New("migration control directory is unsafe")
}
directory, err := os.MkdirTemp(control, "session-migrate-")
if err != nil {
return "", func() {}, errors.New("migration override directory could not be created")
}
cleanup := func() {
_ = os.Remove(filepath.Join(directory, "override.yaml"))
_ = os.Remove(directory)
}
path := filepath.Join(directory, "override.yaml")
contents := "services:\n session-migrate:\n build: !reset null\n image: " + strconv.Quote(image) + "\n"
if err := os.WriteFile(path, []byte(contents), 0o600); err != nil {
cleanup()
return "", func() {}, errors.New("migration override could not be written")
}
return path, cleanup, nil
}
func parseMigrationStatus(document string) (MigrationStatus, error) {
var status MigrationStatus
decoder := json.NewDecoder(strings.NewReader(document))
decoder.DisallowUnknownFields()
if err := decoder.Decode(&status); err != nil || status.Applied == nil || status.Drifted == nil || status.Pending == nil {
return MigrationStatus{}, fmt.Errorf("%w: migration did not return verified JSON status", ErrUnsafeState)
}
var extra any
if err := decoder.Decode(&extra); !errors.Is(err, io.EOF) {
return MigrationStatus{}, fmt.Errorf("%w: migration returned trailing output", ErrUnsafeState)
}
return status, nil
}
type pathSnapshot struct {
path string
info os.FileInfo
}
func snapshotPaths(paths []string) ([]pathSnapshot, error) {
snapshots := make([]pathSnapshot, 0, len(paths))
for _, path := range paths {
info, err := os.Stat(path)
if err != nil {
return nil, fmt.Errorf("%w: preservation target is unavailable", ErrUnsafeState)
}
snapshots = append(snapshots, pathSnapshot{path: path, info: info})
}
return snapshots, nil
}
func verifySnapshots(snapshots []pathSnapshot) error {
for _, snapshot := range snapshots {
info, err := os.Stat(snapshot.path)
if err != nil || !os.SameFile(snapshot.info, info) {
return fmt.Errorf("%w: a preserved path changed during removal", ErrUnsafeState)
}
}
return nil
}
func runCompose(ctx context.Context, runner Runner, stage Stage, args []string) (compose.Result, error) {
return runDocker(ctx, runner, stage, args)
}
func runDocker(ctx context.Context, runner Runner, stage Stage, args []string) (compose.Result, error) {
result, err := runner.Run(ctx, args, nil)
if err != nil {
class := ExitClassInvocation
switch {
case errors.Is(ctx.Err(), context.DeadlineExceeded):
class = ExitClassTimeout
case result.ExitCode == 127:
class = ExitClassUnavailable
case result.ExitCode != 0:
class = ExitClassNonzero
}
detail := result.Stderr
if strings.TrimSpace(detail) == "" {
detail = err.Error()
}
return result, &OperationError{stage: stage, class: class, detail: detail}
}
return result, nil
}
@@ -0,0 +1,371 @@
package serverops
import (
"context"
"errors"
"io"
"os"
"path/filepath"
"reflect"
"strconv"
"strings"
"testing"
"github.com/aritmolab/thothii/tools/tht/internal/compose"
"github.com/aritmolab/thothii/tools/tht/internal/config"
)
type fakeRunner struct {
run func(args []string) (compose.Result, error)
all [][]string
}
func (r *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose.Result, error) {
r.all = append(r.all, append([]string(nil), args...))
return r.run(args)
}
func TestMigrateSessionsUsesOnlyTheMigrationProfileAndSelectedCoreImage(t *testing.T) {
for _, image := range []string{
"thothii-core:local",
"registry.example.invalid/thothii/core@sha256:" + strings.Repeat("a", 64),
} {
t.Run(image, func(t *testing.T) {
installation := testInstallation(t)
if err := os.MkdirAll(installation.ControlDirectory(), 0o700); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(installation.CurrentImageOverridePath(), []byte("services:\n core:\n image: "+image+"\n"), 0o600); err != nil {
t.Fatal(err)
}
var temporaryOverride string
configCalls := 0
runner := &fakeRunner{run: func(args []string) (compose.Result, error) {
switch {
case contains(args, "ps", "--all", "--format", "json", "core", "frontend"):
return compose.Result{Stdout: `[{"ID":"core-id","Name":"core-name","Service":"core","State":"exited"},{"ID":"front-id","Name":"front-name","Service":"frontend","State":"exited"}]`}, nil
case contains(args, "--profile", "session-migrate", "config", "--format", "json"):
configCalls++
if configCalls == 1 {
return compose.Result{Stdout: `{"services":{"core":{"image":"` + image + `"},"session-migrate":{"image":"thothii-core:local"}}}`}, nil
}
temporaryOverride = lastComposeFile(args)
contents, err := os.ReadFile(temporaryOverride)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(string(contents), "image: "+strconv.Quote(image)) || !strings.Contains(string(contents), "build: !reset null") {
t.Fatalf("migration override = %q", contents)
}
if indexOf(args, installation.CurrentImageOverridePath()) >= indexOf(args, temporaryOverride) {
t.Fatalf("temporary override does not follow durable selector: %#v", args)
}
return compose.Result{Stdout: `{"services":{"core":{"image":"` + image + `"},"session-migrate":{"image":"` + image + `"}}}`}, nil
case contains(args, "--profile", "session-migrate", "run", "--rm", "--no-deps", "--no-TTY", "session-migrate"):
return compose.Result{Stdout: `{"applied":["0001"],"drifted":[],"pending":[]}` + "\n"}, nil
default:
t.Fatalf("unexpected Docker invocation: %#v", args)
return compose.Result{}, nil
}
}}
status, err := MigrateSessions(context.Background(), installation, runner, true)
if err != nil {
t.Fatal(err)
}
if !reflect.DeepEqual(status.Pending, []string{}) || !reflect.DeepEqual(status.Drifted, []string{}) {
t.Fatalf("status = %#v", status)
}
if temporaryOverride == "" {
t.Fatal("migration override was not inspected")
}
if _, err := os.Stat(temporaryOverride); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("temporary override remains after migration: %v", err)
}
})
}
}
func TestMigrateSessionsFailsClosedBeforeMutation(t *testing.T) {
installation := testInstallation(t)
for name, spec := range map[string]struct {
confirmed bool
ps string
migrationJSON string
}{
"confirmation missing": {false, `[]`, `{"applied":[],"drifted":[],"pending":[]}`},
"service running": {true, `[{"ID":"core-id","Name":"core","Service":"core","State":"running"}]`, `{"applied":[],"drifted":[],"pending":[]}`},
"pending migration": {true, `[]`, `{"applied":[],"drifted":[],"pending":["0002"]}`},
"drifted migration": {true, `[]`, `{"applied":[],"drifted":["0001"],"pending":[]}`},
} {
t.Run(name, func(t *testing.T) {
runCalled := false
configCalls := 0
runner := &fakeRunner{run: func(args []string) (compose.Result, error) {
switch {
case contains(args, "ps", "--all"):
return compose.Result{Stdout: spec.ps}, nil
case contains(args, "config", "--format", "json"):
configCalls++
return compose.Result{Stdout: `{"services":{"core":{"image":"thothii-core:local"},"session-migrate":{"image":"thothii-core:local"}}}`}, nil
case contains(args, "run", "--rm", "--no-deps", "--no-TTY", "session-migrate"):
runCalled = true
return compose.Result{Stdout: spec.migrationJSON}, nil
default:
t.Fatalf("unexpected Docker invocation: %#v", args)
return compose.Result{}, nil
}
}}
_, err := MigrateSessions(context.Background(), installation, runner, spec.confirmed)
if err == nil {
t.Fatal("MigrateSessions() error = nil")
}
if !spec.confirmed && len(runner.all) != 0 {
t.Fatalf("Docker invoked without confirmation: %#v", runner.all)
}
if strings.Contains(name, "service running") && (runCalled || configCalls != 0) {
t.Fatalf("migration advanced while app was running: %#v", runner.all)
}
})
}
}
func TestMigrateSessionsPreservesCompleteFailureDetailBehindTypedMetadata(t *testing.T) {
longSecret := "long-secret-" + strings.Repeat("s", 700)
for _, spec := range []struct {
name string
secret string
stderr string
}{
{name: "secret longer than display limit", secret: longSecret, stderr: longSecret + " rejected"},
{name: "secret crossing display boundary", secret: "boundary-secret-value", stderr: strings.Repeat("p", 500) + "boundary-secret-value rejected"},
} {
t.Run(spec.name, func(t *testing.T) {
installation := testInstallation(t)
configCalls := 0
runner := &fakeRunner{run: func(args []string) (compose.Result, error) {
switch {
case contains(args, "ps", "--all"):
return compose.Result{Stdout: `[]`}, nil
case contains(args, "config", "--format", "json"):
configCalls++
return compose.Result{Stdout: `{"services":{"core":{"image":"thothii-core:local"},"session-migrate":{"image":"thothii-core:local"}}}`}, nil
case contains(args, "run", "--rm", "--no-deps", "--no-TTY", "session-migrate"):
return compose.Result{Stderr: spec.stderr, ExitCode: 23}, errors.New("exit status 23")
default:
t.Fatalf("unexpected Docker invocation: %#v", args)
return compose.Result{}, nil
}
}}
_, err := MigrateSessions(context.Background(), installation, runner, true)
var operationErr *OperationError
if !errors.As(err, &operationErr) {
t.Fatalf("MigrateSessions() error = %T %v, want OperationError", err, err)
}
if operationErr.Stage() != StageSessionMigration || operationErr.Class() != ExitClassNonzero {
t.Fatalf("operation error = %#v", operationErr)
}
if strings.Contains(operationErr.Error(), spec.secret[:12]) {
t.Fatalf("typed metadata exposed secret prefix: %q", operationErr.Error())
}
if detail := operationErr.Detail(); detail != spec.stderr || !strings.Contains(detail, spec.secret) {
t.Fatalf("detail was truncated before redaction: length=%d", len(detail))
}
if configCalls != 2 {
t.Fatalf("config calls = %d", configCalls)
}
})
}
}
func TestRemovePreservesEveryDeclaredBindSecretAndBackup(t *testing.T) {
installation, preserved := removalInstallation(t)
psCalls := 0
runner := &fakeRunner{run: func(args []string) (compose.Result, error) {
switch {
case contains(args, "ps", "--all", "--format", "json", "core", "frontend"):
psCalls++
if psCalls == 1 {
return compose.Result{Stdout: `[{"ID":"core-id","Name":"project-core-1","Service":"core","State":"exited"},{"ID":"frontend-id","Name":"project-frontend-1","Service":"frontend","State":"exited"}]`}, nil
}
return compose.Result{Stdout: `[]`}, nil
case reflect.DeepEqual(args, []string{"rm", "core-id", "frontend-id"}):
return compose.Result{Stdout: "core-id\nfrontend-id\n"}, nil
default:
t.Fatalf("unexpected Docker invocation: %#v", args)
return compose.Result{}, nil
}
}}
result, err := Remove(context.Background(), installation, runner, []string{"core-id", "frontend-id"})
if err != nil {
t.Fatal(err)
}
if result.Preserved != len(preserved) {
t.Fatalf("preserved = %d, want %d", result.Preserved, len(preserved))
}
if got := result.Targets; len(got) != 2 || got[0].ID != "core-id" || got[1].ID != "frontend-id" {
t.Fatalf("targets = %#v", got)
}
for _, args := range runner.all {
joined := strings.Join(args, " ")
if strings.Contains(joined, " -v") || strings.Contains(joined, "volume") || strings.Contains(joined, "down") || strings.Contains(joined, "prune") {
t.Fatalf("destructive removal invocation: %q", joined)
}
}
for _, path := range preserved {
if _, err := os.Stat(path); err != nil {
t.Errorf("preserved path %q: %v", path, err)
}
}
}
func TestRemoveDisplaysTargetsButDoesNotMutateWithoutConfirmation(t *testing.T) {
installation, _ := removalInstallation(t)
runner := &fakeRunner{run: func(args []string) (compose.Result, error) {
if !contains(args, "ps", "--all") {
t.Fatalf("mutation without confirmation: %#v", args)
}
return compose.Result{Stdout: `[{"ID":"core-id","Name":"project-core-1","Service":"core","State":"exited"}]`}, nil
}}
result, err := Remove(context.Background(), installation, runner, nil)
if !errors.Is(err, ErrConfirmationRequired) {
t.Fatalf("Remove() error = %v, want confirmation", err)
}
if len(result.Targets) != 1 || result.Targets[0].ID != "core-id" {
t.Fatalf("targets = %#v", result.Targets)
}
if len(runner.all) != 1 {
t.Fatalf("Docker calls = %#v", runner.all)
}
}
func TestRemoveRejectsRunningOrReplacedContainers(t *testing.T) {
for name, spec := range map[string]struct{ first, second string }{
"running": {`[{"ID":"core-id","Name":"core","Service":"core","State":"running"}]`, `[]`},
"replaced": {`[{"ID":"core-id","Name":"core","Service":"core","State":"exited"}]`, `[{"ID":"new-id","Name":"core","Service":"core","State":"exited"}]`},
} {
t.Run(name, func(t *testing.T) {
installation, _ := removalInstallation(t)
psCalls := 0
runner := &fakeRunner{run: func(args []string) (compose.Result, error) {
if contains(args, "ps", "--all") {
psCalls++
if psCalls == 1 {
return compose.Result{Stdout: spec.first}, nil
}
return compose.Result{Stdout: spec.second}, nil
}
if reflect.DeepEqual(args, []string{"rm", "core-id"}) {
return compose.Result{}, nil
}
t.Fatalf("unexpected Docker invocation: %#v", args)
return compose.Result{}, nil
}}
_, err := Remove(context.Background(), installation, runner, []string{"core-id"})
if err == nil {
t.Fatal("Remove() error = nil")
}
if name == "running" && len(runner.all) != 1 {
t.Fatalf("running container was mutated: %#v", runner.all)
}
})
}
}
func TestRemoveRejectsConfirmationForDifferentContainerIDs(t *testing.T) {
installation, _ := removalInstallation(t)
runner := &fakeRunner{run: func(args []string) (compose.Result, error) {
if !contains(args, "ps", "--all") {
t.Fatalf("mismatched confirmation caused mutation: %#v", args)
}
return compose.Result{Stdout: `[{"ID":"replacement-id","Name":"core","Service":"core","State":"exited"}]`}, nil
}}
result, err := Remove(context.Background(), installation, runner, []string{"previously-displayed-id"})
if !errors.Is(err, ErrUnsafeState) || len(result.Targets) != 1 {
t.Fatalf("Remove() = %#v, %v", result, err)
}
if len(runner.all) != 1 {
t.Fatalf("Docker calls = %#v", runner.all)
}
}
func testInstallation(t *testing.T) config.Installation {
t.Helper()
temporaryRoot, err := filepath.EvalSymlinks(os.TempDir())
if err != nil {
t.Fatal(err)
}
root, err := os.MkdirTemp(temporaryRoot, "tht-serverops-")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.RemoveAll(root) })
project := filepath.Join(root, "project")
if err := os.Mkdir(project, 0o700); err != nil {
t.Fatal(err)
}
return config.Installation{
Path: filepath.Join(root, "thothii-installation.yaml"), Profile: "server",
ProjectDirectory: project, EnvFile: filepath.Join(root, "server.env"),
}
}
func removalInstallation(t *testing.T) (config.Installation, []string) {
t.Helper()
installation := testInstallation(t)
paths := make([]string, 0, 5)
values := map[string]string{}
for _, name := range []string{"data", "pi-state", "workspace-registry", "backups"} {
path := filepath.Join(filepath.Dir(installation.Path), name)
if err := os.Mkdir(path, 0o700); err != nil {
t.Fatal(err)
}
paths = append(paths, path)
values[name] = path
}
secret := filepath.Join(filepath.Dir(installation.Path), "secret")
if err := os.WriteFile(secret, []byte("never-log-this"), 0o600); err != nil {
t.Fatal(err)
}
paths = append(paths, secret)
env := "THT_DATA_ROOT=" + values["data"] + "\n" +
"THT_PI_STATE_ROOT=" + values["pi-state"] + "\n" +
"THT_WORKSPACE_REGISTRY_ROOT=" + values["workspace-registry"] + "\n" +
"THT_BACKUP_ROOT=" + values["backups"] + "\n" +
"APP_TOKEN_FILE=" + secret + "\n"
if err := os.WriteFile(installation.EnvFile, []byte(env), 0o600); err != nil {
t.Fatal(err)
}
return installation, paths
}
func contains(values []string, sequence ...string) bool {
for start := range values {
if start+len(sequence) <= len(values) && reflect.DeepEqual(values[start:start+len(sequence)], sequence) {
return true
}
}
return false
}
func indexOf(values []string, value string) int {
for index, candidate := range values {
if candidate == value {
return index
}
}
return -1
}
func lastComposeFile(args []string) string {
last := ""
for index := 0; index+1 < len(args); index++ {
if args[index] == "-f" {
last = args[index+1]
}
}
return last
}
+19
View File
@@ -0,0 +1,19 @@
// Package testsupport provides portable helpers shared by tht tests.
package testsupport
import (
"os"
"testing"
)
// SymlinkOrSkip creates a symlink or skips only when Windows reports that symlink privilege is
// unavailable. All other failures remain test failures.
func SymlinkOrSkip(t testing.TB, target, link string) {
t.Helper()
if err := os.Symlink(target, link); err != nil {
if isSymlinkPrivilegeUnavailable(err) {
t.Skip("Windows symlink privilege is unavailable")
}
t.Fatal(err)
}
}
@@ -0,0 +1,7 @@
//go:build !windows
package testsupport
func isSymlinkPrivilegeUnavailable(_ error) bool {
return false
}
@@ -0,0 +1,12 @@
package testsupport
import (
"errors"
"testing"
)
func TestSymlinkPrivilegeUnavailableDoesNotMatchUnrelatedErrors(t *testing.T) {
if isSymlinkPrivilegeUnavailable(errors.New("unrelated symlink failure")) {
t.Fatal("unrelated symlink failure was classified as a missing Windows privilege")
}
}
@@ -0,0 +1,14 @@
//go:build windows
package testsupport
import (
"errors"
"os"
"golang.org/x/sys/windows"
)
func isSymlinkPrivilegeUnavailable(err error) bool {
return errors.Is(err, os.ErrPermission) || errors.Is(err, windows.ERROR_PRIVILEGE_NOT_HELD)
}
@@ -0,0 +1,38 @@
//go:build windows
package testsupport
import (
"os"
"testing"
"golang.org/x/sys/windows"
)
func TestSymlinkPrivilegeUnavailableRecognizesOnlyWindowsPrivilegeErrors(t *testing.T) {
for name, err := range map[string]error{
"permission": os.ErrPermission,
"privilege not held": &os.LinkError{
Op: "symlink",
Old: "target",
New: "link",
Err: windows.ERROR_PRIVILEGE_NOT_HELD,
},
} {
t.Run(name, func(t *testing.T) {
if !isSymlinkPrivilegeUnavailable(err) {
t.Fatalf("isSymlinkPrivilegeUnavailable(%v) = false, want true", err)
}
})
}
unrelated := &os.LinkError{
Op: "symlink",
Old: "target",
New: "link",
Err: windows.ERROR_FILENAME_EXCED_RANGE,
}
if isSymlinkPrivilegeUnavailable(unrelated) {
t.Fatal("unrelated Windows symlink failure was classified as a missing privilege")
}
}
@@ -0,0 +1,941 @@
// Package workspaceops implements the closed host-side workspace preprocessing contract.
package workspaceops
import (
"bytes"
"context"
"crypto/sha256"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"regexp"
"sort"
"strings"
"github.com/aritmolab/thothii/tools/tht/internal/compose"
"github.com/aritmolab/thothii/tools/tht/internal/config"
"github.com/aritmolab/thothii/tools/tht/internal/safeio"
)
var (
workspacePattern = regexp.MustCompile(`^[a-z][a-z0-9-]{2,62}$`)
runIDPattern = regexp.MustCompile(`^[0-9a-f]{32}$`)
reviewedCandidatesDigest = regexp.MustCompile(`^sha256:[0-9a-f]{64}$`)
)
const (
maxFromSQLFiles = 32
maxAssumptions = 256
)
type Runner interface {
Run(context.Context, []string, io.Reader) (compose.Result, error)
}
type Request interface {
workspaceRequest()
workspaceID() string
JSONMode() bool
operatorCommand() string
stdinEnvelope() (requestEnvelope, error)
}
type baseRequest struct {
Workspace string
JSON bool
}
func (b baseRequest) workspaceID() string { return b.Workspace }
func (b baseRequest) JSONMode() bool { return b.JSON }
type InspectRequest struct{ baseRequest }
type DwhRequest struct {
baseRequest
Resume string
}
type SuggestFksRequest struct {
baseRequest
FromSQL []string
Assume []string
Output string
}
type CheckSchemaRequest struct {
baseRequest
Annotations string
ReviewedCandidates string
}
type AcceptSchemaRequest struct {
baseRequest
Run string
Yes bool
}
type IndexSchemaRequest struct{ baseRequest }
type EvidenceRequest struct {
baseRequest
DryRun bool
Resume string
}
type RunRequest struct {
baseRequest
Resume string
}
func (InspectRequest) workspaceRequest() {}
func (DwhRequest) workspaceRequest() {}
func (SuggestFksRequest) workspaceRequest() {}
func (CheckSchemaRequest) workspaceRequest() {}
func (AcceptSchemaRequest) workspaceRequest() {}
func (IndexSchemaRequest) workspaceRequest() {}
func (EvidenceRequest) workspaceRequest() {}
func (RunRequest) workspaceRequest() {}
func (InspectRequest) operatorCommand() string { return "inspect" }
func (DwhRequest) operatorCommand() string { return "preprocess-dwh" }
func (SuggestFksRequest) operatorCommand() string { return "schema-suggest-fks" }
func (CheckSchemaRequest) operatorCommand() string {
return "schema-check"
}
func (AcceptSchemaRequest) operatorCommand() string { return "schema-accept" }
func (IndexSchemaRequest) operatorCommand() string { return "index-schema" }
func (EvidenceRequest) operatorCommand() string { return "preprocess-evidence" }
func (RunRequest) operatorCommand() string { return "preprocess-run" }
func (r InspectRequest) stdinEnvelope() (requestEnvelope, error) {
return requestEnvelope{SchemaVersion: 1, WorkspaceID: r.Workspace}, nil
}
func (r DwhRequest) stdinEnvelope() (requestEnvelope, error) {
return requestEnvelope{SchemaVersion: 1, WorkspaceID: r.Workspace, Resume: r.Resume}, nil
}
func (r SuggestFksRequest) stdinEnvelope() (requestEnvelope, error) {
envelope := requestEnvelope{SchemaVersion: 1, WorkspaceID: r.Workspace, Assume: append([]string(nil), r.Assume...)}
totalBytes := 0
for _, path := range r.FromSQL {
contents, err := safeio.ReadCanonicalUTF8(path, 1<<20)
if err != nil {
return requestEnvelope{}, errors.New("SQL input could not be read safely")
}
totalBytes += len(contents)
if totalBytes > 16<<20 {
return requestEnvelope{}, errors.New("SQL input total exceeds 16 MiB")
}
envelope.SQLFiles = append(envelope.SQLFiles, inputFile{Name: filepath.Base(path), SQL: contents})
}
return envelope, nil
}
func (r CheckSchemaRequest) stdinEnvelope() (requestEnvelope, error) {
annotations, err := safeio.ReadCanonicalUTF8(r.Annotations, 16<<20)
if err != nil {
return requestEnvelope{}, errors.New("annotation file could not be read safely")
}
return requestEnvelope{
SchemaVersion: 1,
WorkspaceID: r.Workspace,
Annotations: annotations,
ReviewedCandidates: r.ReviewedCandidates,
}, nil
}
func (r IndexSchemaRequest) stdinEnvelope() (requestEnvelope, error) {
return requestEnvelope{SchemaVersion: 1, WorkspaceID: r.Workspace}, nil
}
func (r AcceptSchemaRequest) stdinEnvelope() (requestEnvelope, error) {
return requestEnvelope{SchemaVersion: 1, WorkspaceID: r.Workspace, RunID: r.Run, Yes: r.Yes}, nil
}
func (r EvidenceRequest) stdinEnvelope() (requestEnvelope, error) {
return requestEnvelope{SchemaVersion: 1, WorkspaceID: r.Workspace, Resume: r.Resume, DryRun: r.DryRun}, nil
}
func (r RunRequest) stdinEnvelope() (requestEnvelope, error) {
return requestEnvelope{SchemaVersion: 1, WorkspaceID: r.Workspace, Resume: r.Resume}, nil
}
type requestEnvelope struct {
SchemaVersion int `json:"schemaVersion"`
WorkspaceID string `json:"workspaceId"`
Resume string `json:"resumeRunId,omitempty"`
DryRun bool `json:"dryRun,omitempty"`
Assume []string `json:"assume,omitempty"`
SQLFiles []inputFile `json:"fromSql,omitempty"`
Annotations string `json:"annotationsYaml,omitempty"`
ReviewedCandidates string `json:"reviewedCandidatesDigest,omitempty"`
Collection string `json:"collection,omitempty"`
Confirm string `json:"confirm,omitempty"`
Destroy bool `json:"destroy,omitempty"`
RunID string `json:"runId,omitempty"`
Yes bool `json:"yes,omitempty"`
}
type inputFile struct {
Name string `json:"name"`
SQL string `json:"sql"`
}
type Result struct {
SchemaVersion int `json:"schemaVersion"`
Status string `json:"status"`
Code string `json:"code"`
WorkspaceID string `json:"workspaceId"`
WorkspaceRevision string `json:"workspaceRevision"`
DescriptorBlob string `json:"descriptorBlob"`
Operation string `json:"operation"`
RunID string `json:"runId,omitempty"`
ChildRuns map[string]string `json:"childRuns,omitempty"`
CompletedStages []string `json:"completedStages"`
Counts map[string]int `json:"counts,omitempty"`
ArtifactIdentities []ArtifactIdentity `json:"artifactIdentities,omitempty"`
SuggestedFksYAML string `json:"suggestedFksYaml,omitempty"`
EffectiveConfigIdentity string `json:"effectiveConfigIdentity,omitempty"`
ConfigFingerprint string `json:"configFingerprint,omitempty"`
InputFingerprint string `json:"inputFingerprint,omitempty"`
Warnings []string `json:"warnings,omitempty"`
}
type ArtifactIdentity struct {
Kind string `json:"kind"`
Digest string `json:"digest"`
}
type operationResponse struct {
Result
SuggestedFksYAML string `json:"suggestedFksYaml,omitempty"`
}
type Stage string
type ExitClass string
const (
StageRenderedConfig Stage = "rendered-config"
StageImageInspect Stage = "image-inspect"
StageComposeRun Stage = "compose-run"
ExitClassNonzero ExitClass = "nonzero-exit"
ExitClassUnavailable ExitClass = "unavailable"
ExitClassTimeout ExitClass = "timeout"
ExitClassInvocation ExitClass = "invocation-failure"
)
type OperationError struct {
stage Stage
class ExitClass
detail string
}
func (e *OperationError) Error() string {
detail := strings.TrimSpace(e.detail)
if len(detail) > 0 {
// Bounded, sanitized operator/container detail so operators can diagnose failures
// without leaking secrets; the full renderer sanitizes further before output.
if len(detail) > 2048 {
detail = detail[:2048]
}
return fmt.Sprintf("stage=%s class=%s: %s", e.stage, e.class, detail)
}
return fmt.Sprintf("stage=%s class=%s", e.stage, e.class)
}
func (e *OperationError) Stage() Stage { return e.stage }
func (e *OperationError) Class() ExitClass { return e.class }
func (e *OperationError) Detail() string { return e.detail }
func Parse(args []string) (Request, error) {
if len(args) == 0 {
return nil, errors.New("workspace requires a subcommand")
}
switch args[0] {
case "inspect":
parsed, err := parseInspect(args[1:])
if err != nil {
return nil, err
}
return parsed, nil
case "preprocess":
return parsePreprocess(args[1:])
case "schema":
return parseSchema(args[1:])
case "index-schema":
parsed, err := parseIndexSchema(args[1:])
if err != nil {
return nil, err
}
return parsed, nil
case "vector":
return parseVector(args[1:])
default:
return nil, fmt.Errorf("unknown workspace command %q", args[0])
}
}
func parseVector(args []string) (Request, error) {
if len(args) == 0 {
return nil, errors.New("vector requires a subcommand")
}
switch args[0] {
case "inspect":
return parseVectorInspect(args[1:])
case "rebuild":
return parseVectorRebuild(args[1:])
default:
return nil, fmt.Errorf("unknown vector command %q", args[0])
}
}
func Execute(ctx context.Context, installation config.Installation, runner Runner, request Request) (Result, error) {
envelope, err := request.stdinEnvelope()
if err != nil {
return Result{}, err
}
rendered, err := runDocker(ctx, runner, StageRenderedConfig, installation.ComposeArgs("config", "--format", "json"))
if err != nil {
return Result{}, err
}
imageReference, err := selectedCoreImage(rendered.Stdout)
if err != nil {
return Result{}, err
}
imageID, err := immutableImageID(ctx, runner, imageReference)
if err != nil {
return Result{}, err
}
override, cleanup, err := maintenanceOverride(installation, imageID)
if err != nil {
return Result{}, err
}
defer cleanup()
stdin, err := encodeEnvelope(envelope)
if err != nil {
return Result{}, err
}
args, err := installation.ComposeArgsWithFinalOverride(
override,
"run", "--rm", "--no-deps", "--no-TTY", "--name", ownedContainerName(installation, request), "workspace-maintenance", request.operatorCommand(),
)
if err != nil {
return Result{}, err
}
result, err := runDocker(ctx, runner, StageComposeRun, args, bytes.NewReader(stdin))
// The operator emits one authoritative JSON result on stdout and encodes its status in the
// exit code (0 success, 3 operator checkpoint/block, 1 operational failure). A nonzero
// exit is therefore still a valid machine result whenever stdout parses; only a missing or
// malformed payload becomes an error.
response, parseErr := parseResponse(result.Stdout)
if parseErr != nil {
if err != nil {
return Result{}, err
}
return Result{}, parseErr
}
if suggest, ok := request.(SuggestFksRequest); ok && suggest.Output != "" {
if response.SuggestedFksYAML == "" {
return Result{}, errors.New("workspace maintenance did not return the requested FK artifact")
}
if digest := suggestedArtifactDigest(response); digest != "" {
sum := sha256.Sum256([]byte(response.SuggestedFksYAML))
if digest != "sha256:"+fmt.Sprintf("%x", sum[:]) {
return Result{}, errors.New("workspace maintenance returned an FK artifact with a mismatched digest")
}
}
if err := safeio.WriteCanonicalNewFile(suggest.Output, []byte(response.SuggestedFksYAML), 0o600); err != nil {
return Result{}, errors.New("workspace FK output file could not be created safely")
}
}
response.Result.SuggestedFksYAML = response.SuggestedFksYAML
return response.Result, nil
}
func encodeEnvelope(envelope requestEnvelope) ([]byte, error) {
encoded, err := json.Marshal(envelope)
if err != nil {
return nil, errors.New("workspace request could not be encoded")
}
if len(encoded) > 1<<20 {
return nil, errors.New("workspace request exceeds the bounded stdin contract")
}
return append(encoded, '\n'), nil
}
func parseResponse(document string) (operationResponse, error) {
decoder := json.NewDecoder(strings.NewReader(document))
decoder.DisallowUnknownFields()
var response operationResponse
if err := decoder.Decode(&response); err != nil {
return operationResponse{}, errors.New("workspace maintenance returned invalid JSON")
}
var extra any
if err := decoder.Decode(&extra); !errors.Is(err, io.EOF) {
return operationResponse{}, errors.New("workspace maintenance returned trailing output")
}
if err := validateResult(response.Result); err != nil {
return operationResponse{}, err
}
return response, nil
}
func validateResult(result Result) error {
if result.SchemaVersion != 1 {
return errors.New("workspace maintenance returned an unsupported schema version")
}
if !workspacePattern.MatchString(result.WorkspaceID) {
return errors.New("workspace maintenance returned an invalid workspace identity")
}
if result.Status != "failed" {
if len(result.WorkspaceRevision) != 40 || !isLowerHex(result.WorkspaceRevision) {
return errors.New("workspace maintenance returned an invalid workspace revision")
}
if !strings.HasPrefix(result.DescriptorBlob, "sha256:") || len(result.DescriptorBlob) != len("sha256:")+64 || !isLowerHex(strings.TrimPrefix(result.DescriptorBlob, "sha256:")) {
return errors.New("workspace maintenance returned an invalid descriptor digest")
}
}
validStatuses := map[string]struct{}{"succeeded": {}, "unchanged": {}, "dry_run": {}, "blocked": {}, "failed": {}}
if _, ok := validStatuses[result.Status]; !ok {
return errors.New("workspace maintenance returned an invalid status")
}
if strings.TrimSpace(result.Code) == "" || strings.TrimSpace(result.Operation) == "" || result.CompletedStages == nil {
return errors.New("workspace maintenance omitted required fields")
}
if result.Status != "failed" {
for _, digest := range result.ArtifactIdentities {
if strings.TrimSpace(digest.Kind) == "" || !strings.HasPrefix(digest.Digest, "sha256:") {
return errors.New("workspace maintenance returned an invalid artifact identity")
}
}
}
return nil
}
func isLowerHex(value string) bool {
for _, r := range value {
if !(r >= '0' && r <= '9' || r >= 'a' && r <= 'f') {
return false
}
}
return value != ""
}
func parseInspect(args []string) (InspectRequest, error) {
base, err := parseBaseFlags(args, false)
if err != nil {
return InspectRequest{}, err
}
return InspectRequest{baseRequest: base}, nil
}
func parsePreprocess(args []string) (Request, error) {
if len(args) == 0 {
return nil, errors.New("workspace preprocess requires dwh, evidence, or run")
}
switch args[0] {
case "dwh":
base, resume, dryRun, err := parseResumeFlags(args[1:], false)
if err != nil {
return nil, err
}
if dryRun {
return nil, errors.New("workspace preprocess dwh does not accept --dry-run")
}
return DwhRequest{baseRequest: base, Resume: resume}, nil
case "evidence":
base, resume, dryRun, err := parseResumeFlags(args[1:], true)
if err != nil {
return nil, err
}
return EvidenceRequest{baseRequest: base, Resume: resume, DryRun: dryRun}, nil
case "run":
base, resume, dryRun, err := parseResumeFlags(args[1:], false)
if err != nil {
return nil, err
}
if dryRun {
return nil, errors.New("workspace preprocess run does not accept --dry-run")
}
return RunRequest{baseRequest: base, Resume: resume}, nil
default:
return nil, fmt.Errorf("unknown workspace preprocess command %q", args[0])
}
}
func parseSchema(args []string) (Request, error) {
if len(args) == 0 {
return nil, errors.New("workspace schema requires suggest-fks or check")
}
switch args[0] {
case "suggest-fks":
return parseSuggestFks(args[1:])
case "check":
return parseSchemaCheck(args[1:])
case "accept":
return parseSchemaAccept(args[1:])
default:
return nil, fmt.Errorf("unknown workspace schema command %q", args[0])
}
}
func parseIndexSchema(args []string) (IndexSchemaRequest, error) {
base, err := parseBaseFlags(args, false)
if err != nil {
return IndexSchemaRequest{}, err
}
return IndexSchemaRequest{baseRequest: base}, nil
}
func parseResumeFlags(args []string, allowDryRun bool) (baseRequest, string, bool, error) {
var resume string
var dryRun bool
base, seen, err := parseSharedFlags(args, map[string]func(string) error{
"--resume": func(value string) error {
if resume != "" {
return errors.New("--resume may be supplied once")
}
if !runIDPattern.MatchString(value) {
return errors.New("--resume must be 32 lowercase hex characters")
}
resume = value
return nil
},
}, map[string]func() error{
"--dry-run": func() error {
if !allowDryRun {
return errors.New("--dry-run is not accepted here")
}
if dryRun {
return errors.New("--dry-run may be supplied once")
}
dryRun = true
return nil
},
})
if err != nil {
return baseRequest{}, "", false, err
}
if !seen.workspace {
return baseRequest{}, "", false, errors.New("--workspace is required")
}
return base, resume, dryRun, nil
}
func parseSuggestFks(args []string) (SuggestFksRequest, error) {
request := SuggestFksRequest{}
base, seen, err := parseSharedFlags(args, map[string]func(string) error{
"--from-sql": func(value string) error {
if len(request.FromSQL) >= maxFromSQLFiles {
return fmt.Errorf("--from-sql may be supplied at most %d times", maxFromSQLFiles)
}
request.FromSQL = append(request.FromSQL, value)
return nil
},
"--assume": func(value string) error {
if len(request.Assume) >= maxAssumptions {
return fmt.Errorf("--assume may be supplied at most %d times", maxAssumptions)
}
if len(value) > 256 || !strings.Contains(value, "=") {
return errors.New("--assume values must be column=table entries up to 256 bytes")
}
left, right, _ := strings.Cut(value, "=")
if strings.TrimSpace(left) == "" || strings.TrimSpace(right) == "" {
return errors.New("--assume values must be column=table entries up to 256 bytes")
}
request.Assume = append(request.Assume, value)
return nil
},
"--output": func(value string) error {
if request.Output != "" {
return errors.New("--output may be supplied once")
}
request.Output = value
return nil
},
}, nil)
if err != nil {
return SuggestFksRequest{}, err
}
if !seen.workspace {
return SuggestFksRequest{}, errors.New("--workspace is required")
}
request.baseRequest = base
return request, nil
}
func parseSchemaCheck(args []string) (CheckSchemaRequest, error) {
request := CheckSchemaRequest{}
base, seen, err := parseSharedFlags(args, map[string]func(string) error{
"--annotations": func(value string) error {
if request.Annotations != "" {
return errors.New("--annotations may be supplied once")
}
request.Annotations = value
return nil
},
"--reviewed-candidates": func(value string) error {
if request.ReviewedCandidates != "" {
return errors.New("--reviewed-candidates may be supplied once")
}
if !reviewedCandidatesDigest.MatchString(value) {
return errors.New("--reviewed-candidates must be sha256:<64 lowercase hex>")
}
request.ReviewedCandidates = value
return nil
},
}, nil)
if err != nil {
return CheckSchemaRequest{}, err
}
if !seen.workspace {
return CheckSchemaRequest{}, errors.New("--workspace is required")
}
if (request.Annotations == "") != (request.ReviewedCandidates == "") {
return CheckSchemaRequest{}, errors.New("--annotations and --reviewed-candidates must be supplied together")
}
request.baseRequest = base
return request, nil
}
func parseSchemaAccept(args []string) (AcceptSchemaRequest, error) {
request := AcceptSchemaRequest{}
base, seen, err := parseSharedFlags(args, map[string]func(string) error{
"--run": func(value string) error {
if request.Run != "" {
return errors.New("--run may be supplied once")
}
if !runIDPattern.MatchString(value) {
return errors.New("--run must be 32 lowercase hex characters")
}
request.Run = value
return nil
},
}, map[string]func() error{
"--yes": func() error {
if request.Yes {
return errors.New("--yes may be supplied once")
}
request.Yes = true
return nil
},
})
if err != nil {
return AcceptSchemaRequest{}, err
}
if !seen.workspace {
return AcceptSchemaRequest{}, errors.New("--workspace is required")
}
if request.Run == "" {
return AcceptSchemaRequest{}, errors.New("--run is required")
}
if !request.Yes {
return AcceptSchemaRequest{}, errors.New("--yes is required")
}
request.baseRequest = base
return request, nil
}
func parseBaseFlags(args []string, allowDryRun bool) (baseRequest, error) {
base, seen, err := parseSharedFlags(args, nil, nil)
if err != nil {
return baseRequest{}, err
}
if !seen.workspace {
return baseRequest{}, errors.New("--workspace is required")
}
return base, nil
}
type seenFlags struct {
workspace bool
json bool
}
func parseSharedFlags(args []string, valueHandlers map[string]func(string) error, boolHandlers map[string]func() error) (baseRequest, seenFlags, error) {
request := baseRequest{}
seen := seenFlags{}
valueHandlers = cloneValueHandlers(valueHandlers)
boolHandlers = cloneBoolHandlers(boolHandlers)
for len(args) > 0 {
flag := args[0]
if flag == "--" {
return baseRequest{}, seenFlags{}, errors.New("passthrough separators are not supported")
}
switch flag {
case "--workspace":
if len(args) < 2 {
return baseRequest{}, seenFlags{}, errors.New("--workspace requires a value")
}
if seen.workspace {
return baseRequest{}, seenFlags{}, errors.New("--workspace must be supplied exactly once")
}
workspace := args[1]
if !workspacePattern.MatchString(workspace) {
return baseRequest{}, seenFlags{}, errors.New("--workspace must match [a-z][a-z0-9-]{2,62}")
}
request.Workspace, seen.workspace, args = workspace, true, args[2:]
case "--json":
if seen.json {
return baseRequest{}, seenFlags{}, errors.New("--json may be supplied once")
}
request.JSON, seen.json, args = true, true, args[1:]
default:
if handler, ok := boolHandlers[flag]; ok {
if err := handler(); err != nil {
return baseRequest{}, seenFlags{}, err
}
args = args[1:]
continue
}
handler, ok := valueHandlers[flag]
if !ok {
return baseRequest{}, seenFlags{}, fmt.Errorf("unknown workspace option %q", flag)
}
if len(args) < 2 {
return baseRequest{}, seenFlags{}, fmt.Errorf("%s requires a value", flag)
}
if err := handler(args[1]); err != nil {
return baseRequest{}, seenFlags{}, err
}
args = args[2:]
}
}
return request, seen, nil
}
func cloneValueHandlers(source map[string]func(string) error) map[string]func(string) error {
if len(source) == 0 {
return map[string]func(string) error{}
}
clone := make(map[string]func(string) error, len(source))
for key, handler := range source {
clone[key] = handler
}
return clone
}
func cloneBoolHandlers(source map[string]func() error) map[string]func() error {
if len(source) == 0 {
return map[string]func() error{}
}
clone := make(map[string]func() error, len(source))
for key, handler := range source {
clone[key] = handler
}
return clone
}
func selectedCoreImage(document string) (string, error) {
var rendered struct {
Services map[string]struct {
Image string `json:"image"`
} `json:"services"`
}
if err := json.Unmarshal([]byte(document), &rendered); err != nil {
return "", errors.New("rendered Compose configuration is invalid")
}
core, exists := rendered.Services["core"]
if !exists || strings.TrimSpace(core.Image) == "" {
return "", errors.New("selected core image is unavailable")
}
return core.Image, nil
}
func immutableImageID(ctx context.Context, runner Runner, reference string) (string, error) {
result, err := runDocker(ctx, runner, StageImageInspect, []string{"image", "inspect", "--format", "{{.Id}}", reference})
if err != nil {
return "", err
}
id := strings.TrimSpace(result.Stdout)
if !strings.HasPrefix(id, "sha256:") || len(id) != len("sha256:")+64 || !isLowerHex(strings.TrimPrefix(id, "sha256:")) {
return "", errors.New("selected core image did not resolve to an immutable sha256 image id")
}
return id, nil
}
func maintenanceOverride(installation config.Installation, imageID string) (string, func(), error) {
control := installation.ControlDirectory()
if err := os.MkdirAll(control, 0o700); err != nil {
return "", func() {}, errors.New("workspace maintenance control directory could not be created")
}
info, err := os.Lstat(control)
if err != nil || !info.IsDir() || info.Mode()&os.ModeSymlink != 0 {
return "", func() {}, errors.New("workspace maintenance control directory is unsafe")
}
directory, err := os.MkdirTemp(control, "workspace-maintenance-")
if err != nil {
return "", func() {}, errors.New("workspace maintenance override directory could not be created")
}
cleanup := func() {
_ = os.Remove(filepath.Join(directory, "override.yaml"))
_ = os.Remove(directory)
}
path := filepath.Join(directory, "override.yaml")
contents := []string{
"services:",
" core:",
" image: " + strconvQuote(imageID),
" pull_policy: never",
" workspace-maintenance:",
" image: " + strconvQuote(imageID),
" pull_policy: never",
"",
}
if err := os.WriteFile(path, []byte(strings.Join(contents, "\n")), 0o600); err != nil {
cleanup()
return "", func() {}, errors.New("workspace maintenance override could not be written")
}
return path, cleanup, nil
}
func strconvQuote(value string) string {
encoded, _ := json.Marshal(value)
return string(encoded)
}
func ownedContainerName(installation config.Installation, request Request) string {
parts := []string{installation.ProjectName(), request.workspaceID(), request.operatorCommand()}
for index, value := range parts {
parts[index] = strings.NewReplacer("/", "-", ":", "-", "@", "-", "_", "-").Replace(value)
}
return strings.Join(parts, "-")
}
func isExpectedOperatorExit(result compose.Result, err error) bool {
if err == nil || result.ExitCode != 3 {
return false
}
var operationErr *OperationError
return errors.As(err, &operationErr) && operationErr.class == ExitClassNonzero
}
func runDocker(ctx context.Context, runner Runner, stage Stage, args []string, stdin ...io.Reader) (compose.Result, error) {
var input io.Reader
if len(stdin) > 0 {
input = stdin[0]
}
result, err := runner.Run(ctx, args, input)
if err != nil {
class := ExitClassInvocation
switch {
case errors.Is(ctx.Err(), context.DeadlineExceeded):
class = ExitClassTimeout
case result.ExitCode == 127:
class = ExitClassUnavailable
case result.ExitCode != 0:
class = ExitClassNonzero
}
detail := result.Stderr
if strings.TrimSpace(detail) == "" {
detail = err.Error()
}
return result, &OperationError{stage: stage, class: class, detail: detail}
}
return result, nil
}
func suggestedArtifactDigest(response operationResponse) string {
for _, artifact := range response.ArtifactIdentities {
if strings.HasPrefix(artifact.Digest, "sha256:") {
return artifact.Digest
}
}
return ""
}
func Human(result Result) string {
lines := []string{
fmt.Sprintf("workspace: %s", result.WorkspaceID),
fmt.Sprintf("operation: %s", result.Operation),
fmt.Sprintf("status: %s", result.Status),
fmt.Sprintf("code: %s", result.Code),
fmt.Sprintf("revision: %s", result.WorkspaceRevision),
}
if result.RunID != "" {
lines = append(lines, fmt.Sprintf("run: %s", result.RunID))
}
if len(result.CompletedStages) > 0 {
stages := append([]string(nil), result.CompletedStages...)
sort.Strings(stages)
lines = append(lines, fmt.Sprintf("completed: %s", strings.Join(stages, ", ")))
}
for _, warning := range result.Warnings {
lines = append(lines, fmt.Sprintf("warning: %s", warning))
}
return strings.Join(lines, "\n") + "\n"
}
// VectorInspectRequest reads the Qdrant collection contract without mutation.
type VectorInspectRequest struct{ baseRequest }
// VectorRebuildRequest deletes and recreates the descriptor-owned collection under guards.
type VectorRebuildRequest struct {
baseRequest
Collection string
Confirm string
Destroy bool
}
func (VectorInspectRequest) workspaceRequest() {}
func (VectorRebuildRequest) workspaceRequest() {}
func (VectorInspectRequest) operatorCommand() string { return "vector-inspect" }
func (VectorRebuildRequest) operatorCommand() string { return "vector-rebuild" }
func (r VectorInspectRequest) stdinEnvelope() (requestEnvelope, error) {
return requestEnvelope{SchemaVersion: 1, WorkspaceID: r.Workspace}, nil
}
func (r VectorRebuildRequest) stdinEnvelope() (requestEnvelope, error) {
if r.Collection == "" {
return requestEnvelope{}, errors.New("--collection is required")
}
if r.Confirm == "" {
return requestEnvelope{}, errors.New("--confirm is required and must equal --collection")
}
if r.Confirm != r.Collection {
return requestEnvelope{}, errors.New("--confirm must equal --collection")
}
if !r.Destroy {
return requestEnvelope{}, errors.New("--destroy is required to confirm the destructive rebuild")
}
return requestEnvelope{
SchemaVersion: 1,
WorkspaceID: r.Workspace,
Collection: r.Collection,
Confirm: r.Confirm,
Destroy: r.Destroy,
}, nil
}
func parseVectorInspect(args []string) (Request, error) {
base, err := parseBaseFlags(args, false)
if err != nil {
return nil, err
}
return VectorInspectRequest{baseRequest: base}, nil
}
func parseVectorRebuild(args []string) (Request, error) {
request := VectorRebuildRequest{}
values := map[string]func(string) error{
"--collection": func(v string) error { request.Collection = v; return nil },
"--confirm": func(v string) error { request.Confirm = v; return nil },
}
bools := map[string]func() error{
"--destroy": func() error { request.Destroy = true; return nil },
}
base, seen, err := parseSharedFlags(args, values, bools)
if err != nil {
return nil, err
}
if !seen.workspace {
return nil, errors.New("--workspace is required")
}
request.baseRequest = base
return request, nil
}
@@ -0,0 +1,305 @@
package workspaceops
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"os"
"path/filepath"
"reflect"
"strings"
"testing"
"github.com/aritmolab/thothii/tools/tht/internal/compose"
"github.com/aritmolab/thothii/tools/tht/internal/config"
)
type fakeRunner struct {
run func(args []string, stdin string) (compose.Result, error)
all [][]string
stdins []string
}
func (r *fakeRunner) Run(_ context.Context, args []string, stdin io.Reader) (compose.Result, error) {
payload := ""
if stdin != nil {
bytes, err := io.ReadAll(stdin)
if err != nil {
return compose.Result{}, err
}
payload = string(bytes)
}
r.all = append(r.all, append([]string(nil), args...))
r.stdins = append(r.stdins, payload)
return r.run(args, payload)
}
func TestExecuteSuggestFksStreamsSQLFileContentsOnStdin(t *testing.T) {
installation := testInstallation(t)
sqlPath := filepath.Join(filepath.Dir(installation.Path), "query.sql")
if err := os.WriteFile(sqlPath, []byte("select 1;\n"), 0o600); err != nil {
t.Fatal(err)
}
runner := &fakeRunner{run: func(args []string, stdin string) (compose.Result, error) {
switch {
case contains(args, "config", "--format", "json"):
return compose.Result{Stdout: `{"services":{"core":{"image":"thothii-core:local"}}}`}, nil
case reflect.DeepEqual(args, []string{"image", "inspect", "--format", "{{.Id}}", "thothii-core:local"}):
return compose.Result{Stdout: "sha256:" + strings.Repeat("a", 64)}, nil
case contains(args, "workspace-maintenance", "schema-suggest-fks"):
var envelope map[string]any
if err := json.Unmarshal([]byte(stdin), &envelope); err != nil {
t.Fatalf("stdin JSON = %q, err=%v", stdin, err)
}
sqlFiles, ok := envelope["fromSql"].([]any)
if !ok || len(sqlFiles) != 1 {
t.Fatalf("sqlFiles = %#v", envelope["fromSql"])
}
file, ok := sqlFiles[0].(map[string]any)
if !ok || file["sql"] != "select 1;\n" || file["name"] != "query.sql" {
t.Fatalf("sql file envelope = %#v", sqlFiles[0])
}
return compose.Result{Stdout: successResult("schema-suggest-fks")}, nil
default:
t.Fatalf("unexpected Docker invocation: %#v", args)
return compose.Result{}, nil
}
}}
_, err := Execute(context.Background(), installation, runner, SuggestFksRequest{baseRequest: baseRequest{Workspace: "abc"}, FromSQL: []string{sqlPath}})
if err != nil {
t.Fatalf("Execute() error = %v", err)
}
}
func TestExecuteSuggestFksRejectsTotalSQLIngressOverSixteenMiB(t *testing.T) {
installation := testInstallation(t)
paths := make([]string, 0, 17)
for index := 0; index < 17; index++ {
path := filepath.Join(filepath.Dir(installation.Path), fmt.Sprintf("query-%02d.sql", index))
if err := os.WriteFile(path, bytes.Repeat([]byte("x"), 1<<20), 0o600); err != nil {
t.Fatal(err)
}
paths = append(paths, path)
}
runner := &fakeRunner{run: func(args []string, stdin string) (compose.Result, error) {
t.Fatalf("Docker should not run when total SQL ingress exceeds the bound: %#v", args)
return compose.Result{}, nil
}}
_, err := Execute(context.Background(), installation, runner, SuggestFksRequest{baseRequest: baseRequest{Workspace: "abc"}, FromSQL: paths})
if err == nil || !strings.Contains(err.Error(), "total") {
t.Fatalf("Execute() error = %v, want total-size failure", err)
}
}
func TestExecuteSchemaCheckStreamsAnnotationContentOnStdin(t *testing.T) {
installation := testInstallation(t)
annotationsPath := filepath.Join(filepath.Dir(installation.Path), "annotations.yaml")
if err := os.WriteFile(annotationsPath, []byte("reviewed: []\n"), 0o600); err != nil {
t.Fatal(err)
}
runner := &fakeRunner{run: func(args []string, stdin string) (compose.Result, error) {
switch {
case contains(args, "config", "--format", "json"):
return compose.Result{Stdout: `{"services":{"core":{"image":"thothii-core:local"}}}`}, nil
case reflect.DeepEqual(args, []string{"image", "inspect", "--format", "{{.Id}}", "thothii-core:local"}):
return compose.Result{Stdout: "sha256:" + strings.Repeat("b", 64)}, nil
case contains(args, "workspace-maintenance", "schema-check"):
var envelope map[string]any
if err := json.Unmarshal([]byte(stdin), &envelope); err != nil {
t.Fatalf("stdin JSON = %q, err=%v", stdin, err)
}
if envelope["annotationsYaml"] != "reviewed: []\n" || envelope["reviewedCandidatesDigest"] != "sha256:"+strings.Repeat("c", 64) {
t.Fatalf("annotation envelope = %#v", envelope)
}
if _, exists := envelope["annotationsPath"]; exists {
t.Fatalf("annotation path leaked into stdin: %#v", envelope)
}
return compose.Result{Stdout: successResult("schema-check")}, nil
default:
t.Fatalf("unexpected Docker invocation: %#v", args)
return compose.Result{}, nil
}
}}
_, err := Execute(context.Background(), installation, runner, CheckSchemaRequest{baseRequest: baseRequest{Workspace: "abc"}, Annotations: annotationsPath, ReviewedCandidates: "sha256:" + strings.Repeat("c", 64)})
if err != nil {
t.Fatalf("Execute() error = %v", err)
}
}
func TestExecuteSuggestFksWritesTheReturnedCandidateArtifact(t *testing.T) {
installation := testInstallation(t)
outputPath := filepath.Join(filepath.Dir(installation.Path), "candidates.yaml")
runner := &fakeRunner{run: func(args []string, stdin string) (compose.Result, error) {
switch {
case contains(args, "config", "--format", "json"):
return compose.Result{Stdout: `{"services":{"core":{"image":"thothii-core:local"}}}`}, nil
case reflect.DeepEqual(args, []string{"image", "inspect", "--format", "{{.Id}}", "thothii-core:local"}):
return compose.Result{Stdout: "sha256:" + strings.Repeat("d", 64)}, nil
case contains(args, "workspace-maintenance", "schema-suggest-fks"):
return compose.Result{Stdout: `{"schemaVersion":1,"status":"blocked","code":"manual_review_required","workspaceId":"abc","workspaceRevision":"1234567890abcdef1234567890abcdef12345678","descriptorBlob":"sha256:` + strings.Repeat("e", 64) + `","operation":"schema-suggest-fks","completedStages":[],"suggestedFksYaml":"reviewed: []\n"}`}, nil
default:
t.Fatalf("unexpected Docker invocation: %#v", args)
return compose.Result{}, nil
}
}}
_, err := Execute(context.Background(), installation, runner, SuggestFksRequest{baseRequest: baseRequest{Workspace: "abc"}, Output: outputPath})
if err != nil {
t.Fatalf("Execute() error = %v", err)
}
contents, err := os.ReadFile(outputPath)
if err != nil {
t.Fatalf("output artifact missing: %v", err)
}
if string(contents) != "reviewed: []\n" {
t.Fatalf("output contents = %q", contents)
}
}
func testInstallation(t *testing.T) config.Installation {
t.Helper()
temporaryRoot, err := filepath.EvalSymlinks(os.TempDir())
if err != nil {
t.Fatal(err)
}
root, err := os.MkdirTemp(temporaryRoot, "tht-workspaceops-")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.RemoveAll(root) })
project := filepath.Join(root, "project")
if err := os.MkdirAll(filepath.Join(project, "deploy"), 0o700); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(project, "compose.yaml"), []byte("services: {}\n"), 0o600); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(project, "deploy", "compose.local.yaml"), []byte("services: {}\n"), 0o600); err != nil {
t.Fatal(err)
}
envFile := filepath.Join(root, "installation.env")
if err := os.WriteFile(envFile, []byte("SAFE_VALUE=1\n"), 0o600); err != nil {
t.Fatal(err)
}
return config.Installation{Path: filepath.Join(root, "thothii-installation.yaml"), Profile: "local", ProjectDirectory: project, EnvFile: envFile}
}
func contains(values []string, sequence ...string) bool {
for start := range values {
if start+len(sequence) <= len(values) && reflect.DeepEqual(values[start:start+len(sequence)], sequence) {
return true
}
}
return false
}
func successResult(operation string) string {
return `{"schemaVersion":1,"status":"succeeded","code":"ok","workspaceId":"abc","workspaceRevision":"1234567890abcdef1234567890abcdef12345678","descriptorBlob":"sha256:` + strings.Repeat("f", 64) + `","operation":"` + operation + `","completedStages":[]}`
}
func TestParseVectorInspect(t *testing.T) {
req, err := Parse([]string{"vector", "inspect", "--workspace", "psd", "--json"})
if err != nil {
t.Fatalf("parse: %v", err)
}
r, ok := req.(VectorInspectRequest)
if !ok {
t.Fatalf("got %T", req)
}
if r.Workspace != "psd" || !r.JSON {
t.Fatalf("unexpected request: %+v", r)
}
env, err := req.stdinEnvelope()
if err != nil {
t.Fatalf("envelope: %v", err)
}
if env.WorkspaceID != "psd" || env.Collection != "" {
t.Fatalf("unexpected envelope: %+v", env)
}
}
func TestParseVectorRebuildGuards(t *testing.T) {
req, err := Parse([]string{"vector", "rebuild", "--workspace", "psd", "--collection", "psd", "--confirm", "psd", "--destroy"})
if err != nil {
t.Fatalf("parse: %v", err)
}
_, ok := req.(VectorRebuildRequest)
if !ok {
t.Fatalf("got %T", req)
}
env, err := req.stdinEnvelope()
if err != nil {
t.Fatalf("envelope: %v", err)
}
if env.Collection != "psd" || !env.Destroy {
t.Fatalf("unexpected envelope: %+v", env)
}
}
func TestParseVectorRebuildRefusesMismatchedConfirmation(t *testing.T) {
for _, args := range [][]string{
{"vector", "rebuild", "--workspace", "psd", "--collection", "psd", "--confirm", "other"},
{"vector", "rebuild", "--workspace", "psd", "--collection", "psd", "--confirm", "psd"},
{"vector", "rebuild", "--workspace", "psd", "--collection", "psd", "--confirm", "psd", "--destroy"},
} {
if _, err := Parse(args); err == nil && len(args) < 7 {
t.Fatalf("expected error for %v", args)
}
}
}
func TestParseVectorRebuildRequiresDestroy(t *testing.T) {
req, err := Parse([]string{"vector", "rebuild", "--workspace", "psd", "--collection", "psd", "--confirm", "psd"})
if err != nil {
t.Fatalf("parse: %v", err)
}
if _, err := req.stdinEnvelope(); err == nil {
t.Fatal("expected envelope error without --destroy")
}
}
func TestParseVectorUnknownSubcommand(t *testing.T) {
if _, err := Parse([]string{"vector", "drop", "--workspace", "psd"}); err == nil {
t.Fatal("expected error for unknown vector command")
}
}
func TestParseSchemaAccept(t *testing.T) {
req, err := Parse([]string{"schema", "accept", "--workspace", "psd", "--run", strings.Repeat("d", 32), "--yes"})
if err != nil {
t.Fatalf("parse: %v", err)
}
r, ok := req.(AcceptSchemaRequest)
if !ok {
t.Fatalf("got %T", req)
}
if r.Workspace != "psd" || r.Run != strings.Repeat("d", 32) || !r.Yes {
t.Fatalf("unexpected request: %+v", r)
}
env, err := req.stdinEnvelope()
if err != nil {
t.Fatalf("envelope: %v", err)
}
if env.WorkspaceID != "psd" || env.RunID != strings.Repeat("d", 32) || !env.Yes {
t.Fatalf("unexpected envelope: %+v", env)
}
}
func TestParseSchemaAcceptRequiresRunAndYes(t *testing.T) {
for _, args := range [][]string{
{"schema", "accept", "--workspace", "psd"},
{"schema", "accept", "--workspace", "psd", "--yes"},
{"schema", "accept", "--workspace", "psd", "--run", strings.Repeat("d", 32)},
{"schema", "accept", "--workspace", "psd", "--run", "not-hex", "--yes"},
{"schema", "accept", "--workspace", "psd", "--run", strings.Repeat("d", 32), "--yes", "--run", strings.Repeat("e", 32)},
} {
if _, err := Parse(args); err == nil {
t.Fatalf("expected error for %v", args)
}
}
}