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
+750
View File
@@ -0,0 +1,750 @@
// tht is the host-side operator command for a local ThothII installation.
package main
import (
"bufio"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"os/exec"
"path/filepath"
"strconv"
"strings"
"github.com/aritmolab/thothii/tools/tht/internal/compose"
"github.com/aritmolab/thothii/tools/tht/internal/config"
"github.com/aritmolab/thothii/tools/tht/internal/output"
"github.com/aritmolab/thothii/tools/tht/internal/pi"
"github.com/aritmolab/thothii/tools/tht/internal/serverops"
"github.com/aritmolab/thothii/tools/tht/internal/workspaceops"
)
const usage = `Usage: tht --installation <absolute-path>/thothii-installation.yaml <command>
Commands:
status Show the Compose service state.
doctor Validate Docker, Compose, rendered configuration, line endings, volumes, and health.
logs Show the latest 200 sanitized service log lines (bounded; no follow mode).
start Start the installation in the background.
stop Stop the installation.
update --check-only Validate the current installation without changing containers.
sessions migrate --yes
Run only the server session migrator and verify pending=[] and drifted=[].
remove Display exact stopped app container IDs without mutation.
remove --yes ID... Remove only the stopped IDs copied from the preceding display.
pi status Show the Pi version embedded in core.
pi doctor Check Pi preconditions without changing the installation.
pi test Run the temporary Pi/core smoke checks.
pi check Alias for pi test.
pi configure [--provider P --model M --thinking low|medium|high]
Select closed backend defaults interactively on a TTY; all flags are required otherwise.
pi restart --yes [--drain]
Recreate only core with the currently selected Pi image and verify readiness.
pi update --version V --source build --yes [--drain]
Rebuild a pinned Pi version and recreate only core.
pi update --version V --source pull --image IMAGE@sha256:DIGEST --yes [--drain]
Pull an immutable candidate and recreate only core.
pi rollback --yes Restore the image recorded by the latest Pi update.
pi maintenance status
Show the durable core admission-gate state.
pi maintenance recover --yes
Verify a terminal installation, remove stale lifecycle files, and clear maintenance.
pi logs Show the latest 200 sanitized core log lines (bounded; no follow mode).
workspace inspect --workspace ID [--json]
Inspect the active registry snapshot for one workspace.
workspace preprocess dwh --workspace ID [--resume RUN] [--json]
workspace schema suggest-fks --workspace ID [--from-sql FILE]... [--assume COLUMN=TABLE]... [--output FILE] [--json]
workspace schema check --workspace ID [--annotations FILE --reviewed-candidates sha256:HEX] [--json]
workspace schema accept --workspace ID --run RUN --yes [--json]
workspace index-schema --workspace ID [--json]
workspace preprocess evidence --workspace ID [--dry-run] [--resume RUN] [--json]
workspace preprocess run --workspace ID [--resume RUN] [--json]
`
func main() {
os.Exit(run(context.Background(), os.Args[1:], os.Stdout, os.Stderr))
}
func run(ctx context.Context, args []string, stdout, stderr io.Writer) int {
if len(args) == 1 && (args[0] == "--help" || args[0] == "-h") {
fmt.Fprint(stdout, usage)
return 0
}
installationPath, command, commandArgs, err := parseArgs(args)
if err != nil {
fmt.Fprintf(stderr, "tht: %s\n\n%s", err, usage)
return 2
}
installation, err := config.Load(installationPath)
if err != nil {
fmt.Fprintf(stderr, "tht: %s\n", output.Sanitize(err.Error(), nil))
return 2
}
secretFiles, err := installation.SecretFiles()
if err != nil {
fmt.Fprintln(stderr, "tht: installation secret declarations could not be read")
return 2
}
secretValues, err := output.SecretValuesFromFiles(secretFiles)
if err != nil {
fmt.Fprintln(stderr, "tht: declared secret file could not be read")
return 2
}
runner := compose.NewRunner("")
var result compose.Result
switch command {
case "status":
if len(commandArgs) != 0 {
return commandUsageError(stderr, "status does not accept arguments")
}
result, err = runner.Run(ctx, installation.ComposeArgs("ps", "--format", "json"), nil)
case "logs":
logArgs, argumentError := logsArgs(commandArgs)
if argumentError != nil {
return commandUsageError(stderr, argumentError.Error())
}
result, err = runner.Run(ctx, installation.ComposeArgs(logArgs...), nil)
case "start":
if len(commandArgs) != 0 {
return commandUsageError(stderr, "start does not accept arguments")
}
result, err = runner.Run(ctx, installation.ComposeArgs("up", "--detach", "--remove-orphans"), nil)
case "stop":
if len(commandArgs) != 0 {
return commandUsageError(stderr, "stop does not accept arguments")
}
result, err = runner.Run(ctx, installation.ComposeArgs("stop"), nil)
case "update":
if len(commandArgs) != 1 || commandArgs[0] != "--check-only" {
return commandUsageError(stderr, "update currently requires --check-only")
}
result, err = runner.Run(ctx, installation.ComposeArgs("config", "--quiet"), nil)
case "doctor":
if len(commandArgs) != 0 {
return commandUsageError(stderr, "doctor does not accept arguments")
}
return doctor(ctx, installation, runner, secretValues, stdout, stderr)
case "pi":
return piCommand(ctx, installation, runner, commandArgs, secretValues, stdout, stderr)
case "sessions":
if len(commandArgs) != 2 || commandArgs[0] != "migrate" || commandArgs[1] != "--yes" {
return commandUsageError(stderr, "sessions migrate requires --yes")
}
status, operationErr := serverops.MigrateSessions(ctx, installation, runner, true)
if operationErr != nil {
return serverOperationFailure(stderr, operationErr, secretValues)
}
if encodeErr := json.NewEncoder(stdout).Encode(status); encodeErr != nil {
fmt.Fprintln(stderr, "tht: migration status could not be written")
return 1
}
return 0
case "remove":
var confirmedIDs []string
if len(commandArgs) > 0 {
if commandArgs[0] != "--yes" || len(commandArgs) < 2 {
return commandUsageError(stderr, "remove requires either no arguments or --yes followed by every displayed container ID")
}
confirmedIDs = commandArgs[1:]
}
removal, operationErr := serverops.Remove(ctx, installation, runner, confirmedIDs)
writeRemovalTargets(stdout, installation.ProjectName(), removal.Targets)
if errors.Is(operationErr, serverops.ErrConfirmationRequired) {
fmt.Fprint(stderr, "tht: inspect the exact targets above, then re-run with remove --yes")
for _, target := range removal.Targets {
fmt.Fprintf(stderr, " %s", target.ID)
}
fmt.Fprintln(stderr)
return 2
}
if operationErr != nil {
return serverOperationFailure(stderr, operationErr, secretValues)
}
fmt.Fprintf(stdout, "Removed %d stopped app containers; verified %d preserved paths.\n", len(removal.Targets), removal.Preserved)
return 0
case "workspace":
return workspaceCommand(ctx, installation, runner, commandArgs, secretValues, stdout, stderr)
default:
return commandUsageError(stderr, fmt.Sprintf("unknown command %q", command))
}
return writeResult(result, err, secretValues, stdout, stderr)
}
func writeRemovalTargets(outputWriter io.Writer, project string, targets []serverops.Container) {
fmt.Fprintf(outputWriter, "Removal targets for installation project %s:\n", project)
if len(targets) == 0 {
fmt.Fprintln(outputWriter, " (none)")
return
}
for _, target := range targets {
fmt.Fprintf(outputWriter, " service=%s name=%s id=%s state=%s\n", target.Service, target.Name, target.ID, target.State)
}
}
func workspaceCommand(ctx context.Context, installation config.Installation, runner compose.Runner, args []string, secretValues []string, stdout, stderr io.Writer) int {
request, err := workspaceops.Parse(args)
if err != nil {
return commandUsageError(stderr, err.Error())
}
result, err := workspaceops.Execute(ctx, installation, runner, request)
if err != nil {
return workspaceFailure(stderr, err, secretValues)
}
if request.JSONMode() {
encoder := json.NewEncoder(stdout)
encoder.SetEscapeHTML(false)
if encodeErr := encoder.Encode(result); encodeErr != nil {
fmt.Fprintln(stderr, "tht: workspace result could not be written")
return 1
}
} else {
fmt.Fprint(stdout, workspaceops.Human(result))
}
switch result.Status {
case "blocked":
return 3
case "failed":
return 1
default:
return 0
}
}
func workspaceFailure(stderr io.Writer, err error, secretValues []string) int {
message := output.Sanitize(err.Error(), secretValues)
var operationErr *workspaceops.OperationError
if errors.As(err, &operationErr) && operationErr.Detail() != "" {
detail := output.SanitizeDetail(operationErr.Detail(), secretValues)
fmt.Fprintf(stderr, "tht: %s: %s\n", message, detail)
} else {
fmt.Fprintf(stderr, "tht: %s\n", message)
}
return 1
}
func serverOperationFailure(stderr io.Writer, err error, secretValues []string) int {
message := output.Sanitize(err.Error(), secretValues)
var operationErr *serverops.OperationError
if errors.As(err, &operationErr) && operationErr.Detail() != "" {
detail := output.SanitizeDetail(operationErr.Detail(), secretValues)
fmt.Fprintf(stderr, "tht: %s: %s\n", message, detail)
} else {
fmt.Fprintf(stderr, "tht: %s\n", message)
}
if errors.Is(err, serverops.ErrConfirmationRequired) || errors.Is(err, serverops.ErrUnsafeState) {
return 2
}
return 1
}
// installationRunner transforms only Compose invocations into the installation's validated,
// profile-specific argument list. Direct Docker image commands remain host-side and use arguments.
type installationRunner struct {
installation config.Installation
runner compose.Runner
}
func (r installationRunner) SessionInventoryScope() string {
if r.installation.Profile == "local" {
return "mine"
}
return "all"
}
func (r installationRunner) Run(ctx context.Context, args []string, stdin io.Reader) (compose.Result, error) {
if len(args) > 0 && args[0] == "compose" {
return r.runner.Run(ctx, r.installation.ComposeArgs(args[1:]...), stdin)
}
return r.runner.Run(ctx, args, stdin)
}
func piCommand(ctx context.Context, installation config.Installation, runner compose.Runner, args []string, secretValues []string, stdout, stderr io.Writer) int {
if len(args) == 0 {
return commandUsageError(stderr, "pi requires a subcommand")
}
controlled := installationRunner{installation: installation, runner: runner}
switch args[0] {
case "status":
if len(args) != 1 {
return commandUsageError(stderr, "pi status does not accept arguments")
}
version, err := pi.Status(ctx, controlled)
if err != nil {
return piFailure(stderr, err, secretValues)
}
fmt.Fprintf(stdout, "Pi version: %s\n", output.Sanitize(version, secretValues))
return 0
case "doctor":
if len(args) != 1 {
return commandUsageError(stderr, "pi doctor does not accept arguments")
}
if err := pi.Doctor(ctx, controlled); err != nil {
return piFailure(stderr, err, secretValues)
}
fmt.Fprintln(stdout, "Pi preflight checks passed.")
return 0
case "test", "check":
if len(args) != 1 {
return commandUsageError(stderr, "pi test does not accept arguments")
}
if err := pi.Test(ctx, controlled); err != nil {
return piFailure(stderr, err, secretValues)
}
fmt.Fprintln(stdout, "Pi/core smoke checks passed.")
return 0
case "logs":
if len(args) != 1 {
return commandUsageError(stderr, "pi logs does not support --follow; use bounded snapshots")
}
logArgs := []string{"logs", "--tail", "200", "core"}
result, err := controlled.Run(ctx, append([]string{"compose"}, logArgs...), nil)
return writeResult(result, err, secretValues, stdout, stderr)
case "configure":
authFile, authErr := installation.EnvironmentValue("PI_AUTH_FILE")
if authErr != nil || strings.TrimSpace(authFile) == "" {
return commandUsageError(stderr, "PI_AUTH_FILE must name the actual protected host credential file")
}
defaults, err := resolvePiConfigure(ctx, controlled, args[1:], os.Stdin, stdout, stdinIsTTY(os.Stdin))
if err != nil {
return commandUsageError(stderr, err.Error())
}
if err := pi.Configure(ctx, controlled, defaults); err != nil {
return piFailure(stderr, err, secretValues)
}
fmt.Fprintf(stdout, "Pi defaults applied and read back. Provider credentials remain only in the host file %s (mode 0600). Never pass credentials to tht.\n", authFile)
return 0
case "restart":
request, err := parsePiRestartArgs(
args[1:],
installation.RestartStatePath(),
installation.UpdateStatePath(),
)
if err != nil {
return commandUsageError(stderr, err.Error())
}
result, err := pi.Restart(ctx, controlled, request)
if err != nil {
return piFailure(stderr, err, secretValues)
}
fmt.Fprintf(stdout, "Pi core restarted with the existing image; version %s readiness and smoke checks passed.\n", output.Sanitize(result.Version, secretValues))
return 0
case "update":
request, err := parsePiUpdateArgs(
args[1:],
installation.UpdateStatePath(),
installation.RestartStatePath(),
)
if err != nil {
return commandUsageError(stderr, err.Error())
}
result, err := pi.Update(ctx, controlled, request)
if err != nil {
return piFailure(stderr, err, secretValues)
}
if result.Phase == pi.PhaseNoop {
fmt.Fprintf(stdout, "Pi already runs requested version %s; no container was recreated.\n", request.Version)
return 0
}
fmt.Fprintf(stdout, "Pi update verified. Recovery metadata: %s\n", result.StatePath)
return 0
case "rollback":
if len(args) != 2 || args[1] != "--yes" {
return commandUsageError(stderr, "pi rollback requires --yes")
}
result, err := pi.Rollback(
ctx,
controlled,
installation.UpdateStatePath(),
installation.RestartStatePath(),
true,
)
if err != nil {
return piFailure(stderr, err, secretValues)
}
fmt.Fprintf(stdout, "Pi rollback restored the recorded core image. Recovery metadata: %s\n", result.StatePath)
return 0
case "maintenance":
if len(args) == 2 && args[1] == "status" {
status, err := pi.MaintenanceStatus(ctx, controlled)
if err != nil {
return piFailure(stderr, err, secretValues)
}
fmt.Fprintf(stdout, "Pi maintenance active: %t (admissions: %d)\n", status.Active, status.Admissions)
return 0
}
if len(args) == 3 && args[1] == "recover" && args[2] == "--yes" {
if err := pi.RecoverLifecycleMaintenance(
ctx,
controlled,
installation.UpdateStatePath(),
installation.RestartStatePath(),
true,
); err != nil {
return piFailure(stderr, err, secretValues)
}
fmt.Fprintln(stdout, "Pi maintenance recovery verified; stale lifecycle files were removed and admissions are open.")
return 0
}
return commandUsageError(stderr, "pi maintenance requires status or recover --yes")
default:
return commandUsageError(stderr, fmt.Sprintf("unknown pi command %q", args[0]))
}
}
func resolvePiConfigure(
ctx context.Context,
runner pi.Runner,
args []string,
input io.Reader,
prompt io.Writer,
isTTY bool,
) (pi.Defaults, error) {
if len(args) > 0 {
return parsePiConfigureArgs(args)
}
if !isTTY {
return pi.Defaults{}, errors.New("non-interactive pi configure requires --provider --model --thinking")
}
options, err := pi.ConfigurationOptions(ctx, runner)
if err != nil {
return pi.Defaults{}, err
}
providers := uniqueProviders(options)
scanner := bufio.NewScanner(input)
provider, err := numberedChoice(scanner, prompt, "provider", providers)
if err != nil {
return pi.Defaults{}, err
}
models := make([]string, 0)
for _, option := range options {
if option.Provider == provider {
models = append(models, option.ID)
}
}
model, err := numberedChoice(scanner, prompt, "model", models)
if err != nil {
return pi.Defaults{}, err
}
thinking, err := numberedChoice(scanner, prompt, "thinking level", []string{"low", "medium", "high"})
if err != nil {
return pi.Defaults{}, err
}
return pi.Defaults{Provider: provider, Model: model, Thinking: thinking}, nil
}
func uniqueProviders(options []pi.ModelOption) []string {
seen := make(map[string]bool)
providers := make([]string, 0)
for _, option := range options {
if !seen[option.Provider] {
seen[option.Provider] = true
providers = append(providers, option.Provider)
}
}
return providers
}
func numberedChoice(scanner *bufio.Scanner, output io.Writer, label string, choices []string) (string, error) {
if len(choices) == 0 {
return "", fmt.Errorf("Pi returned no %s choices", label)
}
fmt.Fprintf(output, "Select %s:\n", label)
for index, choice := range choices {
fmt.Fprintf(output, " %d) %s\n", index+1, choice)
}
for {
fmt.Fprintf(output, "Choice [1-%d]: ", len(choices))
if !scanner.Scan() {
return "", fmt.Errorf("interactive %s selection ended before a choice was entered", label)
}
selected, err := strconv.Atoi(strings.TrimSpace(scanner.Text()))
if err == nil && selected >= 1 && selected <= len(choices) {
return choices[selected-1], nil
}
fmt.Fprintln(output, "Enter one of the listed numbers.")
}
}
func stdinIsTTY(input *os.File) bool {
info, err := input.Stat()
return err == nil && info.Mode()&os.ModeCharDevice != 0
}
func parsePiConfigureArgs(args []string) (pi.Defaults, error) {
var value pi.Defaults
for len(args) > 0 {
if len(args) < 2 {
return pi.Defaults{}, errors.New("configure options require values")
}
key, v := args[0], args[1]
args = args[2:]
switch key {
case "--provider":
value.Provider = v
case "--model":
value.Model = v
case "--thinking":
value.Thinking = v
default:
return pi.Defaults{}, fmt.Errorf("unknown pi configure option %q", key)
}
}
if value.Provider == "" || value.Model == "" || value.Thinking == "" {
return pi.Defaults{}, errors.New("pi configure requires --provider --model --thinking; THT_LLM_URL stays Compose-managed")
}
return value, nil
}
func parsePiUpdateArgs(args []string, statePath, restartStatePath string) (pi.Request, error) {
request := pi.Request{StatePath: statePath, RestartStatePath: restartStatePath}
for len(args) > 0 {
switch args[0] {
case "--version":
if len(args) < 2 || request.Version != "" {
return pi.Request{}, errors.New("pi update requires one --version <pinned-version>")
}
request.Version, args = args[1], args[2:]
case "--source":
if len(args) < 2 {
return pi.Request{}, errors.New("--source requires build or pull")
}
request.Source, args = pi.Source(args[1]), args[2:]
case "--image":
if len(args) < 2 || request.Image != "" {
return pi.Request{}, errors.New("--image requires one digest-pinned image reference")
}
request.Image, args = args[1], args[2:]
case "--yes":
if request.Confirm {
return pi.Request{}, errors.New("--yes may be supplied once")
}
request.Confirm, args = true, args[1:]
case "--drain":
if request.Drain {
return pi.Request{}, errors.New("--drain may be supplied once")
}
request.Drain, args = true, args[1:]
default:
return pi.Request{}, fmt.Errorf("unknown pi update option %q", args[0])
}
}
if request.Version == "" {
return pi.Request{}, errors.New("pi update requires --version <pinned-version>")
}
if request.Source == "" {
return pi.Request{}, errors.New("pi update requires explicit --source build or pull")
}
if request.Source != pi.BuildSource && request.Source != pi.PullSource {
return pi.Request{}, errors.New("--source requires build or pull")
}
if request.Source == pi.PullSource && request.Image == "" {
return pi.Request{}, errors.New("--source pull requires --image <digest-reference>")
}
if request.Source == pi.BuildSource && request.Image != "" {
return pi.Request{}, errors.New("--image is valid only with --source pull")
}
return request, nil
}
func parsePiRestartArgs(args []string, restartStatePath, updateStatePath string) (pi.RestartRequest, error) {
request := pi.RestartRequest{StatePath: restartStatePath, UpdateStatePath: updateStatePath}
for len(args) > 0 {
switch args[0] {
case "--yes":
if request.Confirm {
return pi.RestartRequest{}, errors.New("--yes may be supplied once")
}
request.Confirm, args = true, args[1:]
case "--drain":
if request.Drain {
return pi.RestartRequest{}, errors.New("--drain may be supplied once")
}
request.Drain, args = true, args[1:]
default:
return pi.RestartRequest{}, errors.New("unknown pi restart option")
}
}
if !request.Confirm {
return pi.RestartRequest{}, errors.New("pi restart requires --yes")
}
return request, nil
}
func piFailure(stderr io.Writer, err error, secretValues []string) int {
code := 1
if errors.Is(err, pi.ErrConfirmationRequired) || errors.Is(err, pi.ErrInvalidRequest) || errors.Is(err, pi.ErrActiveSessions) || errors.Is(err, pi.ErrInterruptedUpdate) || errors.Is(err, pi.ErrInterruptedRestart) {
code = 2
}
var childExit interface{ ExitCode() int }
if errors.As(err, &childExit) && childExit.ExitCode() != 0 {
code = childExit.ExitCode()
}
fmt.Fprintf(stderr, "tht: %s\n", output.Sanitize(err.Error(), secretValues))
return code
}
func parseArgs(args []string) (string, string, []string, error) {
if len(args) < 3 || args[0] != "--installation" {
return "", "", nil, errors.New("--installation <absolute-path> is required before the command")
}
if !filepath.IsAbs(args[1]) {
return "", "", nil, errors.New("--installation must be an absolute path")
}
return args[1], args[2], args[3:], nil
}
func logsArgs(args []string) ([]string, error) {
if len(args) == 0 {
return []string{"logs", "--tail", "200"}, nil
}
return nil, errors.New("logs does not accept arguments; use bounded snapshots")
}
func commandUsageError(stderr io.Writer, message string) int {
fmt.Fprintf(stderr, "tht: %s\n", message)
return 2
}
func writeResult(result compose.Result, err error, secretValues []string, stdout, stderr io.Writer) int {
if result.Stdout != "" {
fmt.Fprint(stdout, output.Sanitize(result.Stdout, secretValues))
}
if result.Stderr != "" {
fmt.Fprint(stderr, output.Sanitize(result.Stderr, secretValues))
}
if err == nil {
return 0
}
if errors.Is(err, exec.ErrNotFound) {
fmt.Fprintln(stderr, "tht: Docker is not installed or is not on PATH")
}
if result.ExitCode != 0 {
return result.ExitCode
}
return 1
}
func doctor(ctx context.Context, installation config.Installation, runner compose.Runner, secretValues []string, stdout, stderr io.Writer) int {
checks := [][]string{
{"version", "--format", "{{.Client.Version}}"},
{"compose", "version", "--short"},
installation.ComposeArgs("config", "--quiet"),
installation.ComposeArgs("config", "--format", "json"),
installation.ComposeArgs("ps", "--format", "json"),
}
var renderedConfig, status string
for index, args := range checks {
result, err := runner.Run(ctx, args, nil)
if err != nil {
return writeResult(result, err, secretValues, stdout, stderr)
}
if index == 3 {
renderedConfig = result.Stdout
}
if index == 4 {
status = result.Stdout
}
}
if err := requireLF(installation.ProjectDirectory); err != nil {
fmt.Fprintf(stderr, "tht: %s\n", err)
return 1
}
if err := requireVolumes(renderedConfig); err != nil {
fmt.Fprintf(stderr, "tht: %s\n", err)
return 1
}
if err := requireHealthyServices(status); err != nil {
fmt.Fprintf(stderr, "tht: %s\n", err)
return 1
}
fmt.Fprintln(stdout, "Doctor checks passed.")
return 0
}
func requireLF(root string) error {
return filepath.WalkDir(root, func(path string, entry os.DirEntry, walkErr error) error {
if walkErr != nil {
return walkErr
}
if entry.IsDir() || entry.Type()&os.ModeSymlink != 0 || !requiresLF(entry.Name()) {
return nil
}
contents, err := os.ReadFile(path)
if err != nil {
return err
}
if strings.Contains(string(contents), "\r\n") {
return fmt.Errorf("CRLF line endings found in %s", filepath.Base(path))
}
return nil
})
}
func requiresLF(name string) bool {
if name == "Dockerfile" || strings.HasPrefix(name, "Dockerfile.") || strings.HasSuffix(name, ".Dockerfile") {
return true
}
for _, suffix := range []string{".sh", ".yml", ".yaml"} {
if strings.HasSuffix(name, suffix) {
return true
}
}
return false
}
func requireVolumes(renderedConfig string) error {
var document struct {
Volumes map[string]json.RawMessage `json:"volumes"`
}
if err := json.Unmarshal([]byte(renderedConfig), &document); err != nil {
return fmt.Errorf("Compose returned invalid rendered configuration")
}
if len(document.Volumes) == 0 {
return errors.New("rendered Compose configuration declares no volumes")
}
return nil
}
func requireHealthyServices(status string) error {
type serviceStatus struct {
Service string `json:"Service"`
State string `json:"State"`
Health string `json:"Health"`
}
var services []serviceStatus
if err := json.Unmarshal([]byte(status), &services); err != nil {
decoder := json.NewDecoder(strings.NewReader(status))
for {
var service serviceStatus
if err := decoder.Decode(&service); errors.Is(err, io.EOF) {
break
} else if err != nil {
return errors.New("Compose returned invalid service status")
}
services = append(services, service)
}
if len(services) == 0 {
return errors.New("Compose returned invalid service status")
}
}
seen := map[string]bool{}
for _, service := range services {
if service.Service != "core" && service.Service != "frontend" {
continue
}
if service.State != "running" || service.Health != "healthy" {
return fmt.Errorf("%s is not healthy", service.Service)
}
seen[service.Service] = true
}
for _, service := range []string{"core", "frontend"} {
if !seen[service] {
return fmt.Errorf("%s service is not running", service)
}
}
return nil
}
File diff suppressed because it is too large Load Diff
+21
View File
@@ -0,0 +1,21 @@
module github.com/aritmolab/thothii/tools/tht
go 1.26.0
toolchain go1.26.5
require gopkg.in/yaml.v3 v3.0.1
require (
github.com/compose-spec/compose-go/v2 v2.14.0
github.com/distribution/reference v0.6.0
github.com/gofrs/flock v0.12.1
github.com/sirupsen/logrus v1.9.1
golang.org/x/sys v0.47.0
)
require (
github.com/kr/text v0.2.0 // indirect
github.com/opencontainers/go-digest v1.0.0 // indirect
github.com/rogpeppe/go-internal v1.15.0 // indirect
)
+39
View File
@@ -0,0 +1,39 @@
github.com/compose-spec/compose-go/v2 v2.14.0 h1:uaJeo5B3+OVlu+Rx2qLBcAdXPEUUzm5nQrRiGJafRAQ=
github.com/compose-spec/compose-go/v2 v2.14.0/go.mod h1:ZU6zlcweCZKyiB7BVfCizQT9XmkEIMFE+PRZydVcsZg=
github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E=
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/distribution/reference v0.6.0 h1:0IXCQ5g4/QMHHkarYzh5l+u8T3t73zM5QvfrDyIgxBk=
github.com/distribution/reference v0.6.0/go.mod h1:BbU0aIcezP1/5jX/8MP0YiH4SdvB5Y4f/wlDRiLyi3E=
github.com/gofrs/flock v0.12.1 h1:MTLVXXHf8ekldpJk3AKicLij9MdwOWkZ+a/jHHZby9E=
github.com/gofrs/flock v0.12.1/go.mod h1:9zxTsyu5xtJ9DK+1tFZyibEV7y3uwDxPPfbxeeHCoD0=
github.com/google/go-cmp v0.5.9 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38=
github.com/google/go-cmp v0.5.9/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE=
github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8Oi/yOhh5U=
github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/rogpeppe/go-internal v1.15.0 h1:D0RCU5rMAp+SpgkiNdrjfJ+LX4J1M32V2NeCY7EJ6hc=
github.com/rogpeppe/go-internal v1.15.0/go.mod h1:DrUVZyrJU+txYW5/1kwtXQSMFio52ZOxX7yM1VHvnxs=
github.com/sirupsen/logrus v1.9.1 h1:Ou41VVR3nMWWmTiEUnj0OlsgOSCUFgsPAOl6jRIcVtQ=
github.com/sirupsen/logrus v1.9.1/go.mod h1:naHLuLoDiP4jHNo9R0sCBMtWGeIprob74mVsIT4qYEQ=
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.9.0 h1:HtqpIVDClZ4nwg75+f6Lvsy/wHu+3BoSGCbBAcpTsTg=
github.com/stretchr/testify v1.9.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
golang.org/x/sys v0.0.0-20220715151400-c0bba94af5f8/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk=
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q=
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
gotest.tools/v3 v3.4.0 h1:ZazjZUfuVeZGLAmlKKuyv3IKP5orXcwtOwDQH6YVr6o=
gotest.tools/v3 v3.4.0/go.mod h1:CtbdzLSsqVhDgMtKsx03ird5YTGB3ar27v0u/yKBW5g=
+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)
}
}
}