diff --git a/backend/src/app.ts b/backend/src/app.ts index 9b5d935f..3aa448b1 100644 --- a/backend/src/app.ts +++ b/backend/src/app.ts @@ -15,6 +15,7 @@ import { settingsRoutes, effectiveSettings } from "./routes/settings.js"; import { createPiModelLister } from "./pi/list-models.js"; import { loadSettings, type Settings } from "./settings/settings-store.js"; import { ReadinessManager } from "./runtime/readiness-manager.js"; +import { createMaintenanceGate } from "./runtime/maintenance-gate.js"; import { WorkspaceRegistry } from "./workspaces/registry.js"; import { createProductionWorkspaceDiagnoser } from "./workspaces/diagnostics.js"; import { workspaceRoutes, type WorkspaceDiagnoser } from "./routes/workspaces.js"; @@ -32,6 +33,8 @@ export interface BuildAppDeps { workspaceRegistry?: WorkspaceRegistry; workspaceDiagnoser?: WorkspaceDiagnoser; workspaceRuntimeSupport?: (workspace: WorkspaceDescriptor) => boolean; + /** Returns true while a host maintenance transaction is preventing new runtimes. */ + maintenanceGate?: () => boolean; } export function buildApp(config: AppConfig, deps?: BuildAppDeps): FastifyInstance { @@ -99,6 +102,7 @@ export function buildApp(config: AppConfig, deps?: BuildAppDeps): FastifyInstanc dwhPrecheck: config.dwhPrecheck, legacyWorkspaceMode: config.legacyWorkspaceMode, workspaceRuntimeSupport, + maintenanceGate: deps?.maintenanceGate ?? createMaintenanceGate(config.maintenanceFile), }); sqlRoutes(app, { tht: tht as ThtRunner, getSettings }); metaRoutes(app, { harnessDir: config.harnessDir, listModels }); diff --git a/backend/src/config.ts b/backend/src/config.ts index 4a663c25..b7f87ec7 100644 --- a/backend/src/config.ts +++ b/backend/src/config.ts @@ -12,6 +12,7 @@ export interface AppConfig { defaults: { provider?: string; model?: string; thinking?: string }; maxPiProcesses: number; settingsFile: string; + maintenanceFile: string; dataRoot?: string; ollamaEnsureTimeoutMs: number; secretsFile?: string; @@ -206,6 +207,7 @@ export function loadConfig(env: Record): AppConfig { defaults: { provider: env.PI_PROVIDER, model: env.PI_MODEL, thinking: env.PI_THINKING }, maxPiProcesses: Number(env.MAX_PI_PROCESSES ?? 4), settingsFile: env.SETTINGS_FILE ?? "data/settings.json", + maintenanceFile: env.THT_MAINTENANCE_FILE ?? "data/maintenance.json", dataRoot: env.THT_DATA_ROOT, ollamaEnsureTimeoutMs: Number(env.OLLAMA_ENSURE_TIMEOUT_MS ?? 60000), secretsFile, diff --git a/backend/src/routes/sessions.ts b/backend/src/routes/sessions.ts index 136af41e..e8086c47 100644 --- a/backend/src/routes/sessions.ts +++ b/backend/src/routes/sessions.ts @@ -37,6 +37,8 @@ export function sessionRoutes( legacyWorkspaceMode?: boolean; /** Fail-closed installation/runtime transport capability check. */ workspaceRuntimeSupport: (workspace: WorkspaceDescriptor) => boolean; + /** Host-controlled admission guard. Existing runtimes deliberately continue. */ + maintenanceGate: () => boolean; }, ) { const lifecycleTails = new Map>(); @@ -75,6 +77,11 @@ export function sessionRoutes( return typeof runner.withPrincipal === "function" ? runner.withPrincipal(principal) : runner; }; + const maintenanceReply = (reply: any) => reply.code(503).send({ + code: "maintenance", + error: "Session admission is temporarily paused for maintenance. Try again shortly.", + }); + /** Include retained historical descriptors so removed workspaces remain resumable. */ const sessionRevisions = async () => { const registry = d.workspaceRegistry as Partial; @@ -272,6 +279,7 @@ export function sessionRoutes( }); app.post("/sessions", async (req, reply) => { + if (d.maintenanceGate()) return maintenanceReply(reply); const b = req.body as { question: string; name?: string; workspace?: string; workspaceId?: string; provider?: string; model?: string; thinking?: string; @@ -502,6 +510,7 @@ export function sessionRoutes( return reply.code(204).send(); }); app.post("/sessions/:id/resume", async (req, reply) => { + if (d.maintenanceGate()) return maintenanceReply(reply); const id = (req.params as any).id; const principal = getPrincipal(req); return withSessionLifecycle(id, async () => { diff --git a/backend/src/runtime/maintenance-gate.ts b/backend/src/runtime/maintenance-gate.ts new file mode 100644 index 00000000..2fa02fa2 --- /dev/null +++ b/backend/src/runtime/maintenance-gate.ts @@ -0,0 +1,10 @@ +import { existsSync } from "node:fs"; + +/** + * A host-side lifecycle transaction creates this marker before it checks for active sessions. + * The marker is intentionally only an admission gate: it must never terminate existing Pi + * processes or make their in-flight work unavailable. + */ +export function createMaintenanceGate(markerFile: string): () => boolean { + return () => existsSync(markerFile); +} diff --git a/backend/src/settings/settings-cli.ts b/backend/src/settings/settings-cli.ts new file mode 100644 index 00000000..87598732 --- /dev/null +++ b/backend/src/settings/settings-cli.ts @@ -0,0 +1,30 @@ +/* Core-side, non-interactive installation-default writer used only through compose exec. + * It accepts no credentials and writes the same SETTINGS_FILE consumed by session creation. */ +import { loadConfig } from "../config.js"; +import { loadSettings, saveSettings, type Settings } from "./settings-store.js"; + +const choice = /^[A-Za-z0-9][A-Za-z0-9._/-]{0,127}$/; + +function value(args: string[], flag: string): string { + const at = args.indexOf(flag); + if (at < 0 || at + 1 >= args.length || args.filter((part) => part === flag).length !== 1) { + throw new Error(`missing ${flag}`); + } + return args[at + 1]; +} + +try { + const args = process.argv.slice(2); + if (args.length !== 6) throw new Error("only provider, model, and thinking may be configured"); + const provider = value(args, "--provider"); + const model = value(args, "--model"); + const thinking = value(args, "--thinking"); + if (!choice.test(provider) || !choice.test(model)) throw new Error("invalid provider or model"); + if (!["low", "medium", "high"].includes(thinking)) throw new Error("invalid thinking level"); + const cfg = loadConfig(process.env); + const next: Settings = { ...loadSettings(cfg), provider, model, thinking }; + saveSettings(cfg, next); +} catch (error) { + process.stderr.write(`settings-cli: ${error instanceof Error ? error.message : "invalid configuration"}\n`); + process.exitCode = 2; +} diff --git a/backend/src/settings/settings-store.ts b/backend/src/settings/settings-store.ts index e09d417e..ba880f0b 100644 --- a/backend/src/settings/settings-store.ts +++ b/backend/src/settings/settings-store.ts @@ -1,4 +1,4 @@ -import { mkdirSync, readFileSync, writeFileSync } from "node:fs"; +import { closeSync, fsyncSync, mkdirSync, openSync, readFileSync, renameSync, writeFileSync } from "node:fs"; import { dirname } from "node:path"; import type { AppConfig } from "../config.js"; @@ -29,6 +29,21 @@ export function loadSettings(cfg: AppConfig): Settings { /** Persist settings (pretty JSON). Creates the parent directory if needed. */ export function saveSettings(cfg: AppConfig, s: Settings): Settings { mkdirSync(dirname(cfg.settingsFile), { recursive: true }); - writeFileSync(cfg.settingsFile, JSON.stringify(s, null, 2) + "\n", "utf8"); + const directory = dirname(cfg.settingsFile); + const temporary = `${cfg.settingsFile}.tmp-${process.pid}-${Date.now()}`; + const fd = openSync(temporary, "wx", 0o600); + try { + writeFileSync(fd, JSON.stringify(s, null, 2) + "\n", "utf8"); + fsyncSync(fd); + } finally { + closeSync(fd); + } + renameSync(temporary, cfg.settingsFile); + // The core image runs Linux. Keep the directory acknowledgement explicit there; Windows + // filesystem replacement semantics are delegated to the host-side Go durable writer. + if (process.platform !== "win32") { + const dirFd = openSync(directory, "r"); + try { fsyncSync(dirFd); } finally { closeSync(dirFd); } + } return s; } diff --git a/backend/test/routes-sessions.test.ts b/backend/test/routes-sessions.test.ts index 5610c95d..3ad07a49 100644 --- a/backend/test/routes-sessions.test.ts +++ b/backend/test/routes-sessions.test.ts @@ -2,7 +2,8 @@ import { test, expect, vi } from "vitest"; import { spawn as nodeSpawn } from "node:child_process"; import path from "node:path"; import os from "node:os"; -import { chmodSync, unlinkSync, writeFileSync } from "node:fs"; +import { chmodSync, unlinkSync, writeFileSync, mkdtempSync, rmSync } from "node:fs"; +import { tmpdir } from "node:os"; import { buildApp as buildRealApp } from "../src/app.js"; import { loadConfig } from "../src/config.js"; import { SseHub } from "../src/sse/sse-hub.js"; @@ -74,6 +75,43 @@ test("upstream requests without a principal fail before a Pi runtime can be crea expect(created).toBe(false); }); +test("maintenance rejects new and resumed session admission without interrupting running sessions", async () => { + const app = buildApp(loadConfig({ AUTH_MODE: "upstream", THT_HARNESS_DIR: "../harness" }), { + maintenanceGate: () => true, + thtRunner: { withPrincipal: () => ({ sessionShow: async () => ({ id: "open", status: "open" }) }) } as any, + }); + + const create = await app.inject({ + method: "POST", url: "/sessions", headers: aliceHeaders, payload: { question: "q" }, + }); + const resume = await app.inject({ method: "POST", url: "/sessions/open/resume", headers: aliceHeaders }); + + expect(create.statusCode).toBe(503); + expect(resume.statusCode).toBe(503); + expect(create.json()).toEqual({ + code: "maintenance", error: "Session admission is temporarily paused for maintenance. Try again shortly.", + }); + expect(resume.json()).toEqual(create.json()); +}); + +test("the on-disk maintenance marker gates admission in an upstream server profile", async () => { + const dir = mkdtempSync(path.join(tmpdir(), "tht-maintenance-")); + try { + const marker = path.join(dir, "maintenance.json"); + writeFileSync(marker, '{"transaction":"test"}\n'); + const app = buildApp(loadConfig({ + AUTH_MODE: "upstream", THT_HARNESS_DIR: "../harness", THT_MAINTENANCE_FILE: marker, + }), { thtRunner: {} as any }); + const response = await app.inject({ + method: "POST", url: "/sessions", headers: aliceHeaders, payload: { question: "q" }, + }); + expect(response.statusCode).toBe(503); + expect(response.json().code).toBe("maintenance"); + } finally { + rmSync(dir, { recursive: true, force: true }); + } +}); + test("session routes conceal foreign or missing sessions and deny SSE before it subscribes", async () => { let subscribed = false; const app = buildApp(loadConfig({ AUTH_MODE: "upstream", THT_HARNESS_DIR: "../harness" }), { diff --git a/compose.yaml b/compose.yaml index cf733104..3eec4da0 100644 --- a/compose.yaml +++ b/compose.yaml @@ -13,6 +13,7 @@ services: THT_BIN: /opt/venv/bin/tht THT_DATA_ROOT: /data SETTINGS_FILE: /data/settings/settings.json + THT_MAINTENANCE_FILE: /data/settings/maintenance.json THT_WORKSPACE_REGISTRY_ROOT: /data/workspace-registry THT_WORKSPACE_GIT_REMOTE: ${THT_WORKSPACE_GIT_REMOTE:?set THT_WORKSPACE_GIT_REMOTE} THT_WORKSPACE_GIT_BRANCH: ${THT_WORKSPACE_GIT_BRANCH:-main} diff --git a/tools/thothctl/cmd/thothctl/main.go b/tools/thothctl/cmd/thothctl/main.go index 3e5ed0fc..e2293da7 100644 --- a/tools/thothctl/cmd/thothctl/main.go +++ b/tools/thothctl/cmd/thothctl/main.go @@ -23,7 +23,7 @@ const usage = `Usage: thothctl --installation /thothii-installati Commands: status Show the Compose service state. doctor Validate Docker, Compose, rendered configuration, line endings, volumes, and health. - logs [--follow] Show sanitized service logs (the default is the latest 200 lines). + logs Show the latest 200 sanitized service log lines. start Start the installation in the background. stop Stop the installation. update --check-only Validate the current installation without changing containers. @@ -31,8 +31,8 @@ Commands: pi doctor Check Pi preconditions without changing the installation. pi test Run the temporary Pi/core smoke checks. pi check Alias for pi test. - pi configure Store non-secret Pi defaults (credentials stay in PI_AUTH_FILE). - pi update Rebuild or pull a pinned Pi image (requires --yes). + pi configure Apply non-secret provider/model/thinking defaults to core (credentials stay in PI_AUTH_FILE). + pi update Rebuild or pull a pinned Pi image (--source build|pull and --yes required). pi rollback --yes Restore the image recorded by the latest Pi update. pi logs Show the latest sanitized core logs. ` @@ -169,10 +169,10 @@ func piCommand(ctx context.Context, installation config.Installation, runner com if err != nil { return commandUsageError(stderr, err.Error()) } - if err := pi.Configure(ctx, controlled, filepath.Join(installation.ProjectDirectory, ".thothctl", "pi-defaults.json"), defaults); err != nil { + if err := pi.Configure(ctx, controlled, defaults); err != nil { return piFailure(stderr, err, secretValues) } - fmt.Fprintln(stdout, "Pi defaults saved. Put credentials only in the configured PI_AUTH_FILE (mode 0600).") + fmt.Fprintln(stdout, "Pi defaults applied and read back. Put credentials only in PI_AUTH_FILE (/home/thoth/.pi/agent/auth.json, mode 0600); never pass credentials to thothctl.") return 0 case "update": request, err := parsePiUpdateArgs(args[1:], filepath.Join(installation.ProjectDirectory, ".thothctl", "update-state.json")) @@ -219,20 +219,18 @@ func parsePiConfigureArgs(args []string) (pi.Defaults, error) { value.Model = v case "--thinking": value.Thinking = v - case "--llm-url": - value.LLMURL = v default: return pi.Defaults{}, fmt.Errorf("unknown pi configure option %q", key) } } - if value.Provider == "" || value.Model == "" || value.Thinking == "" || value.LLMURL == "" { - return pi.Defaults{}, errors.New("pi configure requires --provider --model --thinking --llm-url") + if value.Provider == "" || value.Model == "" || value.Thinking == "" { + return pi.Defaults{}, errors.New("pi configure requires --provider --model --thinking; THT_LLM_URL stays Compose-managed") } return value, nil } func parsePiUpdateArgs(args []string, statePath string) (pi.Request, error) { - request := pi.Request{StatePath: statePath, Source: pi.BuildSource} + request := pi.Request{StatePath: statePath} for len(args) > 0 { switch args[0] { case "--version": @@ -264,15 +262,17 @@ func parsePiUpdateArgs(args []string, statePath string) (pi.Request, error) { return pi.Request{}, fmt.Errorf("unknown pi update option %q", args[0]) } } - if request.Version == "" { - return pi.Request{}, errors.New("pi update requires --version ") - } + if request.Version == "" { return pi.Request{}, errors.New("pi update requires --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 ") } + if request.Source == pi.BuildSource && request.Image != "" { return pi.Request{}, errors.New("--image is valid only with --source pull") } return request, nil } func piFailure(stderr io.Writer, err error, secretValues []string) int { code := 1 - if errors.Is(err, pi.ErrConfirmationRequired) || errors.Is(err, pi.ErrActiveSessions) || errors.Is(err, pi.ErrInterruptedUpdate) { + if errors.Is(err, pi.ErrConfirmationRequired) || errors.Is(err, pi.ErrInvalidRequest) { code = 2 } var childExit interface{ ExitCode() int } diff --git a/tools/thothctl/cmd/thothctl/main_test.go b/tools/thothctl/cmd/thothctl/main_test.go index 64e26d16..3f12f688 100644 --- a/tools/thothctl/cmd/thothctl/main_test.go +++ b/tools/thothctl/cmd/thothctl/main_test.go @@ -368,7 +368,7 @@ func TestRunPiUpdateRequiresExplicitConfirmationWithoutInvokingDocker(t *testing fixture.setEnvironment(t) var stdout, stderr bytes.Buffer - exitCode := run(context.Background(), []string{"--installation", fixture.installationPath, "pi", "update", "--version", "0.81.0"}, &stdout, &stderr) + exitCode := run(context.Background(), []string{"--installation", fixture.installationPath, "pi", "update", "--version", "0.81.0", "--source", "build"}, &stdout, &stderr) if exitCode != 2 { t.Errorf("run() exit code = %d, want 2", exitCode) diff --git a/tools/thothctl/go.mod b/tools/thothctl/go.mod index 72316e51..9c44d308 100644 --- a/tools/thothctl/go.mod +++ b/tools/thothctl/go.mod @@ -6,6 +6,9 @@ require gopkg.in/yaml.v3 v3.0.1 require ( github.com/compose-spec/compose-go/v2 v2.14.0 + github.com/distribution/reference v0.6.0 github.com/sirupsen/logrus v1.9.0 golang.org/x/sys v0.5.0 ) + +require github.com/opencontainers/go-digest v1.0.0 // indirect diff --git a/tools/thothctl/go.sum b/tools/thothctl/go.sum index 593ae01e..b2b7a409 100644 --- a/tools/thothctl/go.sum +++ b/tools/thothctl/go.sum @@ -3,8 +3,12 @@ github.com/compose-spec/compose-go/v2 v2.14.0/go.mod h1:ZU6zlcweCZKyiB7BVfCizQT9 github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/distribution/reference v0.6.0 h1:0IXCQ5g4/QMHHkarYzh5l+u8T3t73zM5QvfrDyIgxBk= +github.com/distribution/reference v0.6.0/go.mod h1:BbU0aIcezP1/5jX/8MP0YiH4SdvB5Y4f/wlDRiLyi3E= github.com/google/go-cmp v0.5.9 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38= github.com/google/go-cmp v0.5.9/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= +github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8Oi/yOhh5U= +github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/sirupsen/logrus v1.9.0 h1:trlNQbNUG3OdDrDil03MCb1H2o9nJ1x4/5LYw7byDE0= diff --git a/tools/thothctl/internal/pi/commands.go b/tools/thothctl/internal/pi/commands.go index 9f5b0a98..650c731e 100644 --- a/tools/thothctl/internal/pi/commands.go +++ b/tools/thothctl/internal/pi/commands.go @@ -7,9 +7,6 @@ import ( "errors" "fmt" "io" - "net/url" - "os" - "path/filepath" "regexp" "strings" @@ -22,22 +19,29 @@ type Defaults struct { Provider string `json:"provider"` Model string `json:"model"` Thinking string `json:"thinking"` - LLMURL string `json:"llm_url"` } -// Configure validates and atomically stores only local non-secret Pi defaults. -func Configure(ctx context.Context, runner Runner, path string, value Defaults) error { +var internalIdentityHeaders = []string{ + "-H", "x-thoth-principal-issuer: thothctl", + "-H", "x-thoth-principal-subject: thothctl-maintenance", + "-H", "x-thoth-principal-display-name: Thothctl maintenance", + "-H", "x-thoth-is-admin: 1", +} + +// Configure changes the backend's real installation settings through a core-side helper. It +// deliberately has no secret or endpoint input: external endpoints remain Compose-owned. +func Configure(ctx context.Context, runner Runner, value Defaults) error { if !choicePattern.MatchString(value.Provider) || !choicePattern.MatchString(value.Model) { return errors.New("provider and model must be supported identifiers") } if value.Thinking != "low" && value.Thinking != "medium" && value.Thinking != "high" { return errors.New("thinking must be low, medium, or high") } - u, err := url.Parse(value.LLMURL) - if err != nil || (u.Scheme != "https" && u.Scheme != "http") || u.Host == "" || u.User != nil || u.RawQuery != "" || u.Fragment != "" { - return errors.New("LLM endpoint must be an http(s) URL without credentials, query, or fragment") - } - models, err := runCompose(ctx, runner, "exec", "-T", "core", "curl", "-fsS", "http://127.0.0.1:8787/models") + before, err := renderedCore(ctx, runner) + if err != nil { return err } + args := append([]string{"exec", "-T", "core", "curl", "-fsS"}, internalIdentityHeaders...) + args = append(args, "http://127.0.0.1:8787/models") + models, err := runCompose(ctx, runner, args...) if err != nil { return commandError("Pi options check", models, err) } @@ -59,39 +63,19 @@ func Configure(ctx context.Context, runner Runner, path string, value Defaults) if !found { return errors.New("provider/model is not in Pi options") } - return writeJSON(path, value) -} - -func writeJSON(path string, value any) error { - contents, err := json.MarshalIndent(value, "", " ") - if err != nil { - return err - } - contents = append(contents, '\n') - if err = os.MkdirAll(filepath.Dir(path), 0o700); err != nil { - return errors.New("could not create Pi configuration directory") - } - temporary, err := os.CreateTemp(filepath.Dir(path), ".pi-defaults-*.tmp") - if err != nil { - return errors.New("could not write Pi configuration") - } - name := temporary.Name() - defer os.Remove(name) - if err = temporary.Chmod(0o600); err == nil { - _, err = temporary.Write(contents) - } - if err == nil { - err = temporary.Sync() - } - if closeErr := temporary.Close(); err == nil { - err = closeErr - } - if err == nil { - err = os.Rename(name, path) - } - if err != nil { - return errors.New("could not atomically write Pi configuration") + result, err := runCompose(ctx, runner, "exec", "-T", "core", "node", "/app/backend/dist/settings/settings-cli.js", "--provider", value.Provider, "--model", value.Model, "--thinking", value.Thinking) + if err != nil { return commandError("Pi installation settings write", result, err) } + settingsArgs := append([]string{"exec", "-T", "core", "curl", "-fsS"}, internalIdentityHeaders...) + settingsArgs = append(settingsArgs, "http://127.0.0.1:8787/settings") + settings, err := runCompose(ctx, runner, settingsArgs...) + if err != nil { return commandError("Pi installation settings read-back", settings, err) } + var saved Defaults + if json.Unmarshal([]byte(settings.Stdout), &saved) != nil || saved.Provider != value.Provider || saved.Model != value.Model || saved.Thinking != value.Thinking { + return errors.New("Pi installation settings read-back did not match requested provider, model, and thinking") } + after, err := renderedCore(ctx, runner) + if err != nil { return err } + if before.ConfigurationSHA != after.ConfigurationSHA { return errors.New("external endpoint configuration changed while configuring Pi") } return nil } @@ -131,7 +115,7 @@ func Doctor(ctx context.Context, runner Runner) error { return commandError("Pi preflight check", result, err) } } - return nil + return Test(ctx, runner) } // Test performs the pre-Task-8 composite smoke through core's private loopback endpoint. @@ -140,7 +124,10 @@ func Test(ctx context.Context, runner Runner) error { return err } for _, path := range []string{"health", "models", "settings"} { - result, err := runCompose(ctx, runner, "exec", "-T", "core", "curl", "-fsS", "http://127.0.0.1:8787/"+path) + args := []string{"exec", "-T", "core", "curl", "-fsS"} + if path != "health" { args = append(args, internalIdentityHeaders...) } + args = append(args, "http://127.0.0.1:8787/"+path) + result, err := runCompose(ctx, runner, args...) if err != nil { return commandError("Pi smoke check", result, err) } @@ -148,13 +135,28 @@ func Test(ctx context.Context, runner Runner) error { if err := json.Unmarshal([]byte(result.Stdout), &payload); err != nil { return fmt.Errorf("Pi smoke check returned invalid %s response", path) } - if _, ok := payload.(map[string]any); !ok { + object, ok := payload.(map[string]any) + if !ok { return fmt.Errorf("Pi smoke check returned invalid %s response", path) } + switch path { + case "health": + if object["status"] != "ok" { return errors.New("Pi smoke health response is not ready") } + case "models": + models, ok := object["models"].([]any) + if !ok || len(models) == 0 { return errors.New("Pi smoke models response is empty") } + valid := false + for _, item := range models { if model, ok := item.(map[string]any); ok && stringField(model, "provider") != "" && stringField(model, "id") != "" { valid = true; break } } + if !valid { return errors.New("Pi smoke models response has no provider/model choices") } + case "settings": + if stringField(object, "provider") == "" || stringField(object, "model") == "" || stringField(object, "thinking") == "" { return errors.New("Pi smoke settings response is incomplete") } + } } return nil } +func stringField(value map[string]any, key string) string { text, _ := value[key].(string); return strings.TrimSpace(text) } + func renderedCore(ctx context.Context, runner Runner) (Image, error) { result, err := runCompose(ctx, runner, "config", "--format", "json") if err != nil { diff --git a/tools/thothctl/internal/pi/commands_test.go b/tools/thothctl/internal/pi/commands_test.go index 90e7ef04..45d99a81 100644 --- a/tools/thothctl/internal/pi/commands_test.go +++ b/tools/thothctl/internal/pi/commands_test.go @@ -2,7 +2,6 @@ package pi import ( "context" - "path/filepath" "strings" "testing" ) @@ -17,16 +16,17 @@ func TestDoctorRequiresExternalEndpointAuthPiStateAndHealth(t *testing.T) { } } -func TestConfigureValidatesBackendModelOptionsAndWritesNoSecrets(t *testing.T) { +func TestConfigureValidatesBackendModelOptionsWritesRealCoreSettingsAndUsesUpstreamIdentity(t *testing.T) { fake := newFakeRunner() - path := filepath.Join(t.TempDir(), "pi-defaults.json") - if err := Configure(context.Background(), fake, path, Defaults{Provider: "provider", Model: "model", Thinking: "medium", LLMURL: "https://llm.example.invalid"}); err != nil { + if err := Configure(context.Background(), fake, Defaults{Provider: "provider", Model: "model", Thinking: "medium"}); err != nil { t.Fatal(err) } - if got := string(readStateBytes(t, path)); strings.Contains(got, "secret") || !strings.Contains(got, "llm.example.invalid") { - t.Fatalf("defaults=%q", got) + assertCalled(t, fake.calls, "node /app/backend/dist/settings/settings-cli.js --provider provider --model model --thinking medium") + assertCalled(t, fake.calls, "x-thoth-principal-subject: thothctl-maintenance") + if got := strings.Join(fake.calls, "\n"); strings.Contains(got, "pi-defaults.json") || strings.Contains(got, "secret") { + t.Fatalf("commands=%q", got) } - if err := Configure(context.Background(), fake, path, Defaults{Provider: "provider", Model: "unknown", Thinking: "medium", LLMURL: "https://llm.example.invalid"}); err == nil { + if err := Configure(context.Background(), fake, Defaults{Provider: "provider", Model: "unknown", Thinking: "medium"}); err == nil { t.Fatal("expected unknown model rejection") } } diff --git a/tools/thothctl/internal/pi/durable_unix.go b/tools/thothctl/internal/pi/durable_unix.go new file mode 100644 index 00000000..8f4c0d31 --- /dev/null +++ b/tools/thothctl/internal/pi/durable_unix.go @@ -0,0 +1,15 @@ +//go:build !windows + +package pi + +import "os" + +// durableReplace acknowledges both the data file and its directory entry. A successful return +// is the strongest atomic replacement guarantee supported by Unix filesystems. +func durableReplace(temporary, target, directory string) error { + if err := os.Rename(temporary, target); err != nil { return err } + dir, err := os.Open(directory) + if err != nil { return err } + defer dir.Close() + return dir.Sync() +} diff --git a/tools/thothctl/internal/pi/durable_windows.go b/tools/thothctl/internal/pi/durable_windows.go new file mode 100644 index 00000000..cd664a3a --- /dev/null +++ b/tools/thothctl/internal/pi/durable_windows.go @@ -0,0 +1,15 @@ +//go:build windows + +package pi + +import "golang.org/x/sys/windows" + +// MoveFileEx requests replacement and write-through on Windows. Directory fsync is not exposed +// by the Windows API in the same form as Unix, so callers must not claim a stronger guarantee. +func durableReplace(temporary, target, _ string) error { + from, err := windows.UTF16PtrFromString(temporary) + if err != nil { return err } + to, err := windows.UTF16PtrFromString(target) + if err != nil { return err } + return windows.MoveFileEx(from, to, windows.MOVEFILE_REPLACE_EXISTING|windows.MOVEFILE_WRITE_THROUGH) +} diff --git a/tools/thothctl/internal/pi/process_unix.go b/tools/thothctl/internal/pi/process_unix.go new file mode 100644 index 00000000..cb158ea3 --- /dev/null +++ b/tools/thothctl/internal/pi/process_unix.go @@ -0,0 +1,16 @@ +//go:build !windows + +package pi + +import ( + "errors" + "os" + "syscall" +) + +func processAlive(pid int) bool { + process, err := os.FindProcess(pid) + if err != nil { return false } + err = process.Signal(syscall.Signal(0)) + return err == nil || errors.Is(err, syscall.EPERM) +} diff --git a/tools/thothctl/internal/pi/process_windows.go b/tools/thothctl/internal/pi/process_windows.go new file mode 100644 index 00000000..8af08940 --- /dev/null +++ b/tools/thothctl/internal/pi/process_windows.go @@ -0,0 +1,14 @@ +//go:build windows + +package pi + +import "golang.org/x/sys/windows" + +func processAlive(pid int) bool { + handle, err := windows.OpenProcess(windows.PROCESS_QUERY_LIMITED_INFORMATION, false, uint32(pid)) + if err != nil { return err == windows.ERROR_ACCESS_DENIED } + defer windows.CloseHandle(handle) + var code uint32 + if windows.GetExitCodeProcess(handle, &code) != nil { return true } + return code == 259 // STILL_ACTIVE +} diff --git a/tools/thothctl/internal/pi/state.go b/tools/thothctl/internal/pi/state.go index 7931c702..3bf89c3c 100644 --- a/tools/thothctl/internal/pi/state.go +++ b/tools/thothctl/internal/pi/state.go @@ -5,13 +5,15 @@ import ( "encoding/json" "errors" "fmt" + "crypto/sha256" "os" "path/filepath" - "runtime" + "sort" + "strings" "time" ) -const stateFileVersion = 1 +const stateFileVersion = 2 // Phase describes the durable point reached by a Pi update. type Phase string @@ -32,6 +34,7 @@ type Image struct { Reference string `json:"reference"` Volumes []string `json:"volumes"` Mounts []Mount `json:"mounts"` + MountFingerprint string `json:"mount_fingerprint"` ConfigurationSHA string `json:"configuration_sha256,omitempty"` } @@ -39,9 +42,10 @@ type Image struct { type Mount struct { Type string `json:"type"` Name string `json:"name,omitempty"` - Source string `json:"source"` + SourceSHA256 string `json:"source_sha256"` Destination string `json:"destination"` RW bool `json:"rw"` + Options string `json:"options,omitempty"` } // Target records the immutable input selected by the operator. Source is either build or a @@ -72,14 +76,14 @@ func readState(path string) (State, error) { if err := json.Unmarshal(contents, &state); err != nil { return State{}, errors.New("update recovery state is invalid") } - if state.Version != stateFileVersion || state.Previous.ID == "" || state.Previous.Reference == "" { + if state.Version != stateFileVersion || state.Previous.ID == "" || state.Previous.Reference == "" || state.Previous.MountFingerprint == "" { return State{}, errors.New("update recovery state is incomplete") } return state, nil } func writeState(path string, state State) error { - if state.Previous.ID == "" || state.Previous.Reference == "" { + if state.Previous.ID == "" || state.Previous.Reference == "" || state.Previous.MountFingerprint == "" { return errors.New("refusing to write incomplete update recovery state") } state.Version = stateFileVersion @@ -89,45 +93,52 @@ func writeState(path string, state State) error { return fmt.Errorf("encode update recovery state: %w", err) } contents = append(contents, '\n') - directory := filepath.Dir(path) - if err := os.MkdirAll(directory, 0o700); err != nil { - return errors.New("could not create update recovery directory") - } - temporary, err := os.CreateTemp(directory, ".update-state-*.tmp") - if err != nil { - return errors.New("could not write update recovery state") - } - temporaryName := temporary.Name() - defer os.Remove(temporaryName) - if err := temporary.Chmod(0o600); err != nil { - temporary.Close() - return errors.New("could not protect update recovery state") - } - if _, err := temporary.Write(contents); err != nil { - temporary.Close() - return errors.New("could not write update recovery state") - } - if err := temporary.Sync(); err != nil { - temporary.Close() - return errors.New("could not durably write update recovery state") - } - if err := temporary.Close(); err != nil { - return errors.New("could not write update recovery state") - } - if err := os.Rename(temporaryName, path); err != nil { - return errors.New("could not finalize update recovery state") - } - if runtime.GOOS != "windows" { - if directoryHandle, err := os.Open(directory); err == nil { - _ = directoryHandle.Sync() - _ = directoryHandle.Close() - } + if err := writeFileDurably(path, ".update-state-", contents); err != nil { + return fmt.Errorf("could not durably write update recovery state: %w", err) } return nil } +func writeFileDurably(path, prefix string, contents []byte) error { + directory := filepath.Dir(path) + if err := os.MkdirAll(directory, 0o700); err != nil { return err } + temporary, err := os.CreateTemp(directory, prefix+"*.tmp") + if err != nil { return err } + temporaryName := temporary.Name() + defer os.Remove(temporaryName) + if err := temporary.Chmod(0o600); err != nil { temporary.Close(); return err } + if _, err := temporary.Write(contents); err != nil { temporary.Close(); return err } + if err := temporary.Sync(); err != nil { temporary.Close(); return err } + if err := temporary.Close(); err != nil { return err } + return durableReplace(temporaryName, path, directory) +} + +func mountSourceHash(source string) string { + sum := sha256.Sum256([]byte(source)) + return fmt.Sprintf("%x", sum[:]) +} + +func mountFingerprint(mounts []Mount) string { + values := make([]string, len(mounts)) + for i, mount := range mounts { + values[i] = strings.Join([]string{mount.Type, mount.Name, mount.SourceSHA256, mount.Destination, fmt.Sprint(mount.RW), mount.Options}, "\x00") + } + sort.Strings(values) + sum := sha256.Sum256([]byte(strings.Join(values, "\n"))) + return fmt.Sprintf("%x", sum[:]) +} + +type lockOwner struct { + PID int `json:"pid"` + Host string `json:"host"` + StartedAt time.Time `json:"started_at"` + Transaction string `json:"transaction"` +} + type updateLock struct{ path string } +var ErrLockHeld = errors.New("another Pi update or rollback is already in progress") + func acquireLock(statePath string) (*updateLock, error) { if err := os.MkdirAll(filepath.Dir(statePath), 0o700); err != nil { return nil, errors.New("could not create Pi update recovery directory") @@ -135,10 +146,36 @@ func acquireLock(statePath string) (*updateLock, error) { path := statePath + ".lock" if err := os.Mkdir(path, 0o700); err != nil { if errors.Is(err, os.ErrExist) { - return nil, errors.New("another Pi update or rollback is already in progress; recovery lock retained") + if reclaimDeadLocalLock(path) { + return acquireLock(statePath) + } + return nil, ErrLockHeld } return nil, errors.New("could not acquire Pi update lock") } + host, err := os.Hostname() + if err != nil { _ = os.Remove(path); return nil, errors.New("could not identify Pi update lock owner") } + owner := lockOwner{PID: os.Getpid(), Host: host, StartedAt: time.Now().UTC(), Transaction: fmt.Sprintf("%d-%d", os.Getpid(), time.Now().UnixNano())} + contents, err := json.Marshal(owner) + if err != nil { _ = os.Remove(path); return nil, errors.New("could not record Pi update lock owner") } + if err := writeFileDurably(filepath.Join(path, "owner.json"), ".owner-", append(contents, '\n')); err != nil { + _ = os.Remove(path) + return nil, errors.New("could not record Pi update lock owner") + } return &updateLock{path: path}, nil } -func (l *updateLock) Release() { _ = os.Remove(l.path) } +func (l *updateLock) Release() { _ = os.Remove(filepath.Join(l.path, "owner.json")); _ = os.Remove(l.path) } + +// reclaimDeadLocalLock is deliberately conservative: a malformed, remote, or merely old lock +// is recovery-required. Only a process we can prove is gone on this machine is reclaimed. +func reclaimDeadLocalLock(path string) bool { + contents, err := os.ReadFile(filepath.Join(path, "owner.json")) + if err != nil { return false } + var owner lockOwner + if json.Unmarshal(contents, &owner) != nil || owner.PID <= 0 || owner.Host == "" { return false } + host, err := os.Hostname() + if err != nil || owner.Host != host { return false } + if processAlive(owner.PID) { return false } + if err := os.Remove(filepath.Join(path, "owner.json")); err != nil { return false } + return os.Remove(path) == nil +} diff --git a/tools/thothctl/internal/pi/update.go b/tools/thothctl/internal/pi/update.go index 42b293c6..f410c66b 100644 --- a/tools/thothctl/internal/pi/update.go +++ b/tools/thothctl/internal/pi/update.go @@ -10,14 +10,16 @@ import ( "sort" "strings" "time" + + "github.com/distribution/reference" ) var ( ErrConfirmationRequired = errors.New("update requires --yes after reviewing the planned Pi version") ErrActiveSessions = errors.New("active sessions must be drained before updating Pi; use --drain only after they are complete") ErrInterruptedUpdate = errors.New("a previous Pi update is incomplete; run pi rollback --yes before starting another update") + ErrInvalidRequest = errors.New("invalid Pi lifecycle request") versionPattern = regexp.MustCompile(`^[0-9]+(?:\.[0-9]+){1,3}(?:[-+][0-9A-Za-z.-]+)?$`) - digestPattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._/:@-]*@sha256:[a-f0-9]{64}$`) ) // Source chooses whether the candidate is built from this checkout or pulled from an immutable image. @@ -45,7 +47,7 @@ type Result struct { } // Update performs a recoverable core-only Pi update using the default Compose command layout. -func Update(ctx context.Context, runner Runner, request Request) (Result, error) { +func Update(ctx context.Context, runner Runner, request Request) (result Result, retErr error) { lock, err := acquireLock(request.StatePath) if err != nil { return Result{StatePath: request.StatePath}, err @@ -58,42 +60,39 @@ func Update(ctx context.Context, runner Runner, request Request) (Result, error) return Result{StatePath: request.StatePath}, ErrConfirmationRequired } if !versionPattern.MatchString(request.Version) { - return Result{StatePath: request.StatePath}, errors.New("Pi version must be an explicit pinned version") - } - if request.Source == "" { - request.Source = BuildSource + return Result{StatePath: request.StatePath}, fmt.Errorf("%w: Pi version must be an explicit pinned version", ErrInvalidRequest) } + if request.Source == "" { return Result{StatePath: request.StatePath}, fmt.Errorf("%w: Pi update requires an explicit source: build or pull", ErrInvalidRequest) } if request.Source != BuildSource && request.Source != PullSource { - return Result{StatePath: request.StatePath}, errors.New("Pi update source must be build or pull") + return Result{StatePath: request.StatePath}, fmt.Errorf("%w: Pi update source must be build or pull", ErrInvalidRequest) } - if request.Source == PullSource && !digestPattern.MatchString(request.Image) { - return Result{StatePath: request.StatePath}, errors.New("pulled Pi image must use an immutable sha256 digest") + if request.Source == PullSource { + canonical, err := canonicalDigestReference(request.Image) + if err != nil { return Result{StatePath: request.StatePath}, fmt.Errorf("%w: %v", ErrInvalidRequest, err) } + request.Image = canonical } if old, err := readState(request.StatePath); err == nil && old.Phase != PhaseVerified && old.Phase != PhaseRolledBack && old.Phase != PhaseNoop { return Result{StatePath: request.StatePath}, ErrInterruptedUpdate } else if err != nil && !errors.Is(err, os.ErrNotExist) { return Result{StatePath: request.StatePath}, err } + if err := setMaintenance(ctx, runner, true); err != nil { return Result{StatePath: request.StatePath}, err } + defer func() { + if clearErr := setMaintenance(context.Background(), runner, false); clearErr != nil { + result = Result{Phase: PhaseFailed, StatePath: request.StatePath} + if retErr == nil { retErr = errors.New("maintenance admission gate could not be cleared: recovery required") + } else { retErr = fmt.Errorf("%w; maintenance admission gate could not be cleared: recovery required", retErr) } + } + }() running, err := activeSessions(ctx, runner) if err != nil { return Result{StatePath: request.StatePath}, err } - frontendStopped := false if running { if !request.Drain { return Result{StatePath: request.StatePath}, ErrActiveSessions } - stopped, stopErr := runCompose(ctx, runner, "stop", "frontend") - if stopErr != nil { - return Result{StatePath: request.StatePath}, commandError("frontend admission gate", stopped, stopErr) - } - frontendStopped = true - defer func() { - if frontendStopped { - _, _ = runCompose(context.Background(), runner, "up", "--detach", "frontend") - } - }() for attempts := 0; attempts < 30; attempts++ { running, err = activeSessions(ctx, runner) if err != nil { @@ -134,41 +133,34 @@ func Update(ctx context.Context, runner Runner, request Request) (Result, error) } state.Phase = PhaseBuilding - if err := writeState(request.StatePath, state); err != nil { - return Result{StatePath: request.StatePath}, err - } + if err := writeState(request.StatePath, state); err != nil { return Result{Phase: PhaseFailed, StatePath: request.StatePath}, err } if err := prepareCandidate(ctx, runner, request, previous.Reference); err != nil { - state.Phase, state.Error = PhaseFailed, "candidate image preparation failed" - _ = writeState(request.StatePath, state) - return Result{Phase: PhaseFailed, StatePath: request.StatePath}, err + return compensate(ctx, runner, request.StatePath, state, err) } - if frontendStopped { + if request.Drain { running, err = activeSessions(ctx, runner) if err != nil { - return Result{Phase: PhaseFailed, StatePath: request.StatePath}, err + return compensate(ctx, runner, request.StatePath, state, err) } if running { - state.Phase, state.Error = PhaseFailed, "new session admitted while draining" - _ = writeState(request.StatePath, state) - return Result{Phase: PhaseFailed, StatePath: request.StatePath}, ErrActiveSessions + return compensate(ctx, runner, request.StatePath, state, ErrActiveSessions) } } if err := recreateCore(ctx, runner); err != nil { - state.Phase, state.Error = PhaseFailed, "core recreation failed" - _ = writeState(request.StatePath, state) - return Result{Phase: PhaseFailed, StatePath: request.StatePath}, err + return compensate(ctx, runner, request.StatePath, state, err) } state.Phase = PhaseRecreated - state.Candidate, _ = runningImage(ctx, runner, previous.Reference) + state.Candidate, err = runningImage(ctx, runner, previous.Reference) + if err != nil { return compensate(ctx, runner, request.StatePath, state, err) } if err := writeState(request.StatePath, state); err != nil { - return Result{Phase: PhaseRecreated, StatePath: request.StatePath}, err + return compensate(ctx, runner, request.StatePath, state, err) } if err := verifyCandidate(ctx, runner, request.Version, previous); err != nil { - return rollbackAfterFailure(ctx, runner, request.StatePath, state, err) + return compensate(ctx, runner, request.StatePath, state, err) } state.Phase, state.Error = PhaseVerified, "" if err := writeState(request.StatePath, state); err != nil { - return Result{Phase: PhaseVerified, StatePath: request.StatePath}, err + return compensate(ctx, runner, request.StatePath, state, err) } return Result{Phase: PhaseVerified, StatePath: request.StatePath}, nil } @@ -189,25 +181,35 @@ func Rollback(ctx context.Context, runner Runner, statePath string, confirm bool } if err := restore(ctx, runner, state.Previous); err != nil { state.Phase, state.Error = PhaseFailed, "rollback failed" - _ = writeState(statePath, state) + if writeErr := writeState(statePath, state); writeErr != nil { return Result{Phase: PhaseFailed, StatePath: statePath}, errors.New("rollback failed and recovery state could not be persisted") } return Result{Phase: PhaseFailed, StatePath: statePath}, err } state.Phase, state.Error = PhaseRolledBack, "" if err := writeState(statePath, state); err != nil { - return Result{Phase: PhaseRolledBack, StatePath: statePath}, err + return Result{Phase: PhaseFailed, StatePath: statePath}, errors.New("rollback restored the core but recovery state could not be persisted") } return Result{Phase: PhaseRolledBack, StatePath: statePath}, nil } -func rollbackAfterFailure(ctx context.Context, runner Runner, statePath string, state State, cause error) (Result, error) { +func compensate(ctx context.Context, runner Runner, statePath string, state State, cause error) (Result, error) { if restoreErr := restore(ctx, runner, state.Previous); restoreErr != nil { state.Phase, state.Error = PhaseFailed, "candidate verification and automatic rollback failed" - _ = writeState(statePath, state) - return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("candidate verification failed; automatic rollback also failed") + if writeErr := writeState(statePath, state); writeErr != nil { + return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("update failed and rollback proof failed; recovery state could not be persisted") + } + return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("update failed; automatic rollback also failed: recovery required") } state.Phase, state.Error = PhaseRolledBack, "" - _ = writeState(statePath, state) - return Result{Phase: PhaseRolledBack, StatePath: statePath}, fmt.Errorf("candidate verification failed; previous core image was restored") + if writeErr := writeState(statePath, state); writeErr != nil { + return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("previous core image was restored but recovery state write failed: recovery required") + } + return Result{Phase: PhaseRolledBack, StatePath: statePath}, fmt.Errorf("update failed; previous core image was restored") +} + +func recordFailure(path string, state State, label string, cause error) error { + state.Phase, state.Error = PhaseFailed, label + if err := writeState(path, state); err != nil { return fmt.Errorf("%w; recovery state write failed", cause) } + return cause } func sourceValue(request Request) string { @@ -217,6 +219,29 @@ func sourceValue(request Request) string { return string(BuildSource) } +func canonicalDigestReference(value string) (string, error) { + if strings.Contains(value, "://") || strings.ContainsAny(value, "?#") || strings.Contains(value, "@") && strings.Contains(strings.Split(value, "@")[0], ":") && strings.Contains(strings.Split(value, "@")[0], "//") { + return "", errors.New("pulled Pi image must be a credential-free canonical sha256 digest reference") + } + parsed, err := reference.ParseAnyReference(value) + if err != nil { return "", errors.New("pulled Pi image must be a valid canonical sha256 digest reference") } + canonical, ok := parsed.(reference.Canonical) + if !ok || canonical.Digest().Algorithm().String() != "sha256" || len(canonical.Digest().Encoded()) != 64 { + return "", errors.New("pulled Pi image must use an immutable sha256 digest") + } + return reference.FamiliarString(canonical), nil +} + +// The command text is fixed; no operator input or host path is interpolated into the core shell. +// The marker lives alongside SETTINGS_FILE's named/bind-mounted directory and is read by backend. +func setMaintenance(ctx context.Context, runner Runner, enabled bool) error { + command := "mkdir -p /data/settings && : > /data/settings/maintenance.json && chmod 600 /data/settings/maintenance.json" + if !enabled { command = "rm -f /data/settings/maintenance.json" } + result, err := runCompose(ctx, runner, "exec", "-T", "core", "sh", "-ceu", command) + if err != nil { return commandError("maintenance admission gate", result, err) } + return nil +} + func activeSessions(ctx context.Context, runner Runner) (bool, error) { result, err := runCompose(ctx, runner, "exec", "-T", "core", "tht", "session", "list", "--json") if err != nil { @@ -251,20 +276,27 @@ func runningImage(ctx context.Context, runner Runner, reference string) (Image, if err != nil { return Image{}, commandError("core volume check", mounts, err) } - var contract []Mount - if err := json.Unmarshal([]byte(mounts.Stdout), &contract); err != nil { + var raw []struct { + Type string `json:"Type"`; Name string `json:"Name"`; Source string `json:"Source"` + Destination string `json:"Destination"`; RW bool `json:"RW"`; Mode string `json:"Mode"` + Propagation string `json:"Propagation"`; Driver string `json:"Driver"` + } + if err := json.Unmarshal([]byte(mounts.Stdout), &raw); err != nil { return Image{}, errors.New("core returned invalid persistence mount data") } - if len(contract) == 0 { + if len(raw) == 0 { return Image{}, errors.New("core has no persistence mounts to preserve") } - volumes := make([]string, 0, len(contract)) - for _, mount := range contract { + contract := make([]Mount, 0, len(raw)) + volumes := make([]string, 0, len(raw)) + for _, mount := range raw { + if mount.Type == "" || mount.Source == "" || mount.Destination == "" { return Image{}, errors.New("core returned incomplete persistence mount data") } + contract = append(contract, Mount{Type: mount.Type, Name: mount.Name, SourceSHA256: mountSourceHash(mount.Source), Destination: mount.Destination, RW: mount.RW, Options: strings.Join([]string{mount.Mode, mount.Propagation, mount.Driver}, "\x00")}) if mount.Type == "volume" && mount.Name != "" { volumes = append(volumes, mount.Name) } } - return Image{ID: strings.TrimSpace(image.Stdout), Reference: reference, Volumes: volumes, Mounts: contract}, nil + return Image{ID: strings.TrimSpace(image.Stdout), Reference: reference, Volumes: volumes, Mounts: contract, MountFingerprint: mountFingerprint(contract)}, nil } func prepareCandidate(ctx context.Context, runner Runner, request Request, reference string) error { @@ -382,7 +414,7 @@ func sameMounts(left, right []Mount) bool { return false } key := func(m Mount) string { - return m.Type + "\x00" + m.Name + "\x00" + m.Source + "\x00" + m.Destination + "\x00" + fmt.Sprint(m.RW) + return m.Type + "\x00" + m.Name + "\x00" + m.SourceSHA256 + "\x00" + m.Destination + "\x00" + fmt.Sprint(m.RW) + "\x00" + m.Options } a, b := make([]string, len(left)), make([]string, len(right)) for i := range left { diff --git a/tools/thothctl/internal/pi/update_test.go b/tools/thothctl/internal/pi/update_test.go index 054b1ee0..c7e9e901 100644 --- a/tools/thothctl/internal/pi/update_test.go +++ b/tools/thothctl/internal/pi/update_test.go @@ -3,6 +3,7 @@ package pi import ( "context" "errors" + "fmt" "io" "os" "path/filepath" @@ -103,10 +104,9 @@ func TestUpdateDoesNotRecreateWhenPreflightOrBuildFails(t *testing.T) { if err == nil { t.Fatal("Update() error = nil, want failure") } - if result.Phase == PhaseRolledBack { - t.Fatalf("pre-recreate failure unexpectedly rolled back: %+v", result) - } - assertNotCalled(t, fake.calls, "force-recreate") + if failure == "preflight" && result.Phase == PhaseRolledBack { t.Fatalf("preflight failure unexpectedly rolled back: %+v", result) } + if failure == "build" && result.Phase != PhaseRolledBack { t.Fatalf("candidate build failure must compensate: %+v", result) } + if failure == "preflight" { assertNotCalled(t, fake.calls, "force-recreate") } }) } } @@ -162,7 +162,7 @@ func TestRollbackRestoresInterruptedOrPreviouslyRecordedState(t *testing.T) { func TestUpdateRefusesToOverwriteInterruptedRecoveryState(t *testing.T) { fake := newFakeRunner() statePath := filepath.Join(t.TempDir(), "state.json") - writeStateForTest(t, statePath, State{Phase: PhaseRecreated, Previous: Image{ID: "sha256:old", Reference: "thothii-core:local", Volumes: []string{"settings"}}}) + writeStateForTest(t, statePath, State{Phase: PhaseRecreated, Previous: Image{ID: "sha256:old", Reference: "thothii-core:local", Volumes: []string{"settings"}, MountFingerprint: "recorded"}}) _, err := Update(context.Background(), fake, Request{StatePath: statePath, Version: "0.81.0", Source: BuildSource, Confirm: true}) if !errors.Is(err, ErrInterruptedUpdate) { t.Fatalf("Update() error = %v, want interrupted update error", err) @@ -180,6 +180,17 @@ func TestRunningImageCapturesServerBindAndNamedMountIdentity(t *testing.T) { if len(image.Mounts) != 3 || image.Mounts[0].Type != "bind" || image.Mounts[0].Destination != "/data" { t.Fatalf("mounts = %#v", image.Mounts) } + if strings.Contains(fmt.Sprint(image), "/srv/thothii") || image.Mounts[0].SourceSHA256 == "" || image.MountFingerprint == "" { + t.Fatalf("mount contract leaked a server source or lacks a safe fingerprint: %#v", image) + } +} + +func TestCanonicalDigestReferenceRejectsCredentialsAndURLForms(t *testing.T) { + valid := "registry.example.invalid/thothii-core@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" + if got, err := canonicalDigestReference(valid); err != nil || got != valid { t.Fatalf("canonicalDigestReference() = %q, %v", got, err) } + for _, invalid := range []string{"https://registry.example.invalid/a@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", "user:pass@registry.example/a@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", "registry.example/a@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa?token=x"} { + if _, err := canonicalDigestReference(invalid); err == nil { t.Fatalf("accepted unsafe reference %q", invalid) } + } } type fakeRunner struct { @@ -188,12 +199,13 @@ type fakeRunner struct { version string activeSessions bool built bool + currentImage string volumes []string mountsJSON string } func newFakeRunner() *fakeRunner { - return &fakeRunner{version: "0.80.3", volumes: []string{"settings", "pi-state", "sessions", "workspace-registry"}} + return &fakeRunner{version: "0.80.3", currentImage: "sha256:old", volumes: []string{"settings", "pi-state", "sessions", "workspace-registry"}} } func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose.Result, error) { @@ -201,6 +213,7 @@ func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose f.calls = append(f.calls, call) if strings.Contains(call, "image tag sha256:old") { f.fail = "" + f.currentImage = "sha256:old" } if f.fail == "preflight" && strings.Contains(call, "config --format json") { return compose.Result{ExitCode: 1}, errors.New("provider token=secret") @@ -214,7 +227,7 @@ func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose if f.fail == "version" && f.built && strings.Contains(call, "pi --version") && strings.Contains(call, "exec") { return compose.Result{ExitCode: 1}, errors.New("version token=secret") } - if f.fail == "smoke" && strings.Contains(call, "curl -fsS http://127.0.0.1:8787/models") { + if f.fail == "smoke" && f.built && strings.Contains(call, "127.0.0.1:8787/models") { return compose.Result{ExitCode: 1}, errors.New("smoke token=secret") } switch { @@ -223,7 +236,7 @@ func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose case strings.Contains(call, "ps -q core"): return compose.Result{Stdout: "core-container\n"}, nil case strings.Contains(call, "inspect --format {{.Image}}"): - return compose.Result{Stdout: "sha256:old\n"}, nil + return compose.Result{Stdout: f.currentImage + "\n"}, nil case strings.Contains(call, "inspect --format {{json .Mounts}}"): if f.mountsJSON != "" { return compose.Result{Stdout: f.mountsJSON}, nil @@ -238,6 +251,7 @@ func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose case strings.Contains(call, "compose build"): f.built = true f.version = "0.81.0" + f.currentImage = "sha256:candidate" return compose.Result{}, nil case strings.Contains(call, "pi --version"): return compose.Result{Stdout: f.version + "\n"}, nil