fix: harden pi maintenance lifecycle

This commit is contained in:
2026-08-04 19:09:04 +02:00
parent 8fde1f81c7
commit 0b9ad7f53f
21 changed files with 432 additions and 171 deletions
+4
View File
@@ -15,6 +15,7 @@ import { settingsRoutes, effectiveSettings } from "./routes/settings.js";
import { createPiModelLister } from "./pi/list-models.js"; import { createPiModelLister } from "./pi/list-models.js";
import { loadSettings, type Settings } from "./settings/settings-store.js"; import { loadSettings, type Settings } from "./settings/settings-store.js";
import { ReadinessManager } from "./runtime/readiness-manager.js"; import { ReadinessManager } from "./runtime/readiness-manager.js";
import { createMaintenanceGate } from "./runtime/maintenance-gate.js";
import { WorkspaceRegistry } from "./workspaces/registry.js"; import { WorkspaceRegistry } from "./workspaces/registry.js";
import { createProductionWorkspaceDiagnoser } from "./workspaces/diagnostics.js"; import { createProductionWorkspaceDiagnoser } from "./workspaces/diagnostics.js";
import { workspaceRoutes, type WorkspaceDiagnoser } from "./routes/workspaces.js"; import { workspaceRoutes, type WorkspaceDiagnoser } from "./routes/workspaces.js";
@@ -32,6 +33,8 @@ export interface BuildAppDeps {
workspaceRegistry?: WorkspaceRegistry; workspaceRegistry?: WorkspaceRegistry;
workspaceDiagnoser?: WorkspaceDiagnoser; workspaceDiagnoser?: WorkspaceDiagnoser;
workspaceRuntimeSupport?: (workspace: WorkspaceDescriptor) => boolean; 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 { export function buildApp(config: AppConfig, deps?: BuildAppDeps): FastifyInstance {
@@ -99,6 +102,7 @@ export function buildApp(config: AppConfig, deps?: BuildAppDeps): FastifyInstanc
dwhPrecheck: config.dwhPrecheck, dwhPrecheck: config.dwhPrecheck,
legacyWorkspaceMode: config.legacyWorkspaceMode, legacyWorkspaceMode: config.legacyWorkspaceMode,
workspaceRuntimeSupport, workspaceRuntimeSupport,
maintenanceGate: deps?.maintenanceGate ?? createMaintenanceGate(config.maintenanceFile),
}); });
sqlRoutes(app, { tht: tht as ThtRunner, getSettings }); sqlRoutes(app, { tht: tht as ThtRunner, getSettings });
metaRoutes(app, { harnessDir: config.harnessDir, listModels }); metaRoutes(app, { harnessDir: config.harnessDir, listModels });
+2
View File
@@ -12,6 +12,7 @@ export interface AppConfig {
defaults: { provider?: string; model?: string; thinking?: string }; defaults: { provider?: string; model?: string; thinking?: string };
maxPiProcesses: number; maxPiProcesses: number;
settingsFile: string; settingsFile: string;
maintenanceFile: string;
dataRoot?: string; dataRoot?: string;
ollamaEnsureTimeoutMs: number; ollamaEnsureTimeoutMs: number;
secretsFile?: string; secretsFile?: string;
@@ -206,6 +207,7 @@ export function loadConfig(env: Record<string, string | undefined>): AppConfig {
defaults: { provider: env.PI_PROVIDER, model: env.PI_MODEL, thinking: env.PI_THINKING }, defaults: { provider: env.PI_PROVIDER, model: env.PI_MODEL, thinking: env.PI_THINKING },
maxPiProcesses: Number(env.MAX_PI_PROCESSES ?? 4), maxPiProcesses: Number(env.MAX_PI_PROCESSES ?? 4),
settingsFile: env.SETTINGS_FILE ?? "data/settings.json", settingsFile: env.SETTINGS_FILE ?? "data/settings.json",
maintenanceFile: env.THT_MAINTENANCE_FILE ?? "data/maintenance.json",
dataRoot: env.THT_DATA_ROOT, dataRoot: env.THT_DATA_ROOT,
ollamaEnsureTimeoutMs: Number(env.OLLAMA_ENSURE_TIMEOUT_MS ?? 60000), ollamaEnsureTimeoutMs: Number(env.OLLAMA_ENSURE_TIMEOUT_MS ?? 60000),
secretsFile, secretsFile,
+9
View File
@@ -37,6 +37,8 @@ export function sessionRoutes(
legacyWorkspaceMode?: boolean; legacyWorkspaceMode?: boolean;
/** Fail-closed installation/runtime transport capability check. */ /** Fail-closed installation/runtime transport capability check. */
workspaceRuntimeSupport: (workspace: WorkspaceDescriptor) => boolean; workspaceRuntimeSupport: (workspace: WorkspaceDescriptor) => boolean;
/** Host-controlled admission guard. Existing runtimes deliberately continue. */
maintenanceGate: () => boolean;
}, },
) { ) {
const lifecycleTails = new Map<string, Promise<void>>(); const lifecycleTails = new Map<string, Promise<void>>();
@@ -75,6 +77,11 @@ export function sessionRoutes(
return typeof runner.withPrincipal === "function" ? runner.withPrincipal(principal) : runner; 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. */ /** Include retained historical descriptors so removed workspaces remain resumable. */
const sessionRevisions = async () => { const sessionRevisions = async () => {
const registry = d.workspaceRegistry as Partial<WorkspaceRegistry>; const registry = d.workspaceRegistry as Partial<WorkspaceRegistry>;
@@ -272,6 +279,7 @@ export function sessionRoutes(
}); });
app.post("/sessions", async (req, reply) => { app.post("/sessions", async (req, reply) => {
if (d.maintenanceGate()) return maintenanceReply(reply);
const b = req.body as { const b = req.body as {
question: string; name?: string; workspace?: string; workspaceId?: string; question: string; name?: string; workspace?: string; workspaceId?: string;
provider?: string; model?: string; thinking?: string; provider?: string; model?: string; thinking?: string;
@@ -502,6 +510,7 @@ export function sessionRoutes(
return reply.code(204).send(); return reply.code(204).send();
}); });
app.post("/sessions/:id/resume", async (req, reply) => { app.post("/sessions/:id/resume", async (req, reply) => {
if (d.maintenanceGate()) return maintenanceReply(reply);
const id = (req.params as any).id; const id = (req.params as any).id;
const principal = getPrincipal(req); const principal = getPrincipal(req);
return withSessionLifecycle(id, async () => { return withSessionLifecycle(id, async () => {
+10
View File
@@ -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);
}
+30
View File
@@ -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;
}
+17 -2
View File
@@ -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 { dirname } from "node:path";
import type { AppConfig } from "../config.js"; 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. */ /** Persist settings (pretty JSON). Creates the parent directory if needed. */
export function saveSettings(cfg: AppConfig, s: Settings): Settings { export function saveSettings(cfg: AppConfig, s: Settings): Settings {
mkdirSync(dirname(cfg.settingsFile), { recursive: true }); 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; return s;
} }
+39 -1
View File
@@ -2,7 +2,8 @@ import { test, expect, vi } from "vitest";
import { spawn as nodeSpawn } from "node:child_process"; import { spawn as nodeSpawn } from "node:child_process";
import path from "node:path"; import path from "node:path";
import os from "node:os"; 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 { buildApp as buildRealApp } from "../src/app.js";
import { loadConfig } from "../src/config.js"; import { loadConfig } from "../src/config.js";
import { SseHub } from "../src/sse/sse-hub.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); 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 () => { test("session routes conceal foreign or missing sessions and deny SSE before it subscribes", async () => {
let subscribed = false; let subscribed = false;
const app = buildApp(loadConfig({ AUTH_MODE: "upstream", THT_HARNESS_DIR: "../harness" }), { const app = buildApp(loadConfig({ AUTH_MODE: "upstream", THT_HARNESS_DIR: "../harness" }), {
+1
View File
@@ -13,6 +13,7 @@ services:
THT_BIN: /opt/venv/bin/tht THT_BIN: /opt/venv/bin/tht
THT_DATA_ROOT: /data THT_DATA_ROOT: /data
SETTINGS_FILE: /data/settings/settings.json SETTINGS_FILE: /data/settings/settings.json
THT_MAINTENANCE_FILE: /data/settings/maintenance.json
THT_WORKSPACE_REGISTRY_ROOT: /data/workspace-registry THT_WORKSPACE_REGISTRY_ROOT: /data/workspace-registry
THT_WORKSPACE_GIT_REMOTE: ${THT_WORKSPACE_GIT_REMOTE:?set THT_WORKSPACE_GIT_REMOTE} THT_WORKSPACE_GIT_REMOTE: ${THT_WORKSPACE_GIT_REMOTE:?set THT_WORKSPACE_GIT_REMOTE}
THT_WORKSPACE_GIT_BRANCH: ${THT_WORKSPACE_GIT_BRANCH:-main} THT_WORKSPACE_GIT_BRANCH: ${THT_WORKSPACE_GIT_BRANCH:-main}
+14 -14
View File
@@ -23,7 +23,7 @@ const usage = `Usage: thothctl --installation <absolute-path>/thothii-installati
Commands: Commands:
status Show the Compose service state. status Show the Compose service state.
doctor Validate Docker, Compose, rendered configuration, line endings, volumes, and health. 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. start Start the installation in the background.
stop Stop the installation. stop Stop the installation.
update --check-only Validate the current installation without changing containers. 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 doctor Check Pi preconditions without changing the installation.
pi test Run the temporary Pi/core smoke checks. pi test Run the temporary Pi/core smoke checks.
pi check Alias for pi test. pi check Alias for pi test.
pi configure Store non-secret Pi defaults (credentials stay in PI_AUTH_FILE). 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 (requires --yes). 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 rollback --yes Restore the image recorded by the latest Pi update.
pi logs Show the latest sanitized core logs. 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 { if err != nil {
return commandUsageError(stderr, err.Error()) 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) 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 return 0
case "update": case "update":
request, err := parsePiUpdateArgs(args[1:], filepath.Join(installation.ProjectDirectory, ".thothctl", "update-state.json")) 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 value.Model = v
case "--thinking": case "--thinking":
value.Thinking = v value.Thinking = v
case "--llm-url":
value.LLMURL = v
default: default:
return pi.Defaults{}, fmt.Errorf("unknown pi configure option %q", key) return pi.Defaults{}, fmt.Errorf("unknown pi configure option %q", key)
} }
} }
if value.Provider == "" || value.Model == "" || value.Thinking == "" || value.LLMURL == "" { if value.Provider == "" || value.Model == "" || value.Thinking == "" {
return pi.Defaults{}, errors.New("pi configure requires --provider --model --thinking --llm-url") return pi.Defaults{}, errors.New("pi configure requires --provider --model --thinking; THT_LLM_URL stays Compose-managed")
} }
return value, nil return value, nil
} }
func parsePiUpdateArgs(args []string, statePath string) (pi.Request, error) { 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 { for len(args) > 0 {
switch args[0] { switch args[0] {
case "--version": 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]) return pi.Request{}, fmt.Errorf("unknown pi update option %q", args[0])
} }
} }
if request.Version == "" { if request.Version == "" { return pi.Request{}, errors.New("pi update requires --version <pinned-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 return request, nil
} }
func piFailure(stderr io.Writer, err error, secretValues []string) int { func piFailure(stderr io.Writer, err error, secretValues []string) int {
code := 1 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 code = 2
} }
var childExit interface{ ExitCode() int } var childExit interface{ ExitCode() int }
+1 -1
View File
@@ -368,7 +368,7 @@ func TestRunPiUpdateRequiresExplicitConfirmationWithoutInvokingDocker(t *testing
fixture.setEnvironment(t) fixture.setEnvironment(t)
var stdout, stderr bytes.Buffer 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 { if exitCode != 2 {
t.Errorf("run() exit code = %d, want 2", exitCode) t.Errorf("run() exit code = %d, want 2", exitCode)
+3
View File
@@ -6,6 +6,9 @@ require gopkg.in/yaml.v3 v3.0.1
require ( require (
github.com/compose-spec/compose-go/v2 v2.14.0 github.com/compose-spec/compose-go/v2 v2.14.0
github.com/distribution/reference v0.6.0
github.com/sirupsen/logrus v1.9.0 github.com/sirupsen/logrus v1.9.0
golang.org/x/sys v0.5.0 golang.org/x/sys v0.5.0
) )
require github.com/opencontainers/go-digest v1.0.0 // indirect
+4
View File
@@ -3,8 +3,12 @@ github.com/compose-spec/compose-go/v2 v2.14.0/go.mod h1:ZU6zlcweCZKyiB7BVfCizQT9
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.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 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= 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 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38=
github.com/google/go-cmp v0.5.9/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= 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 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/sirupsen/logrus v1.9.0 h1:trlNQbNUG3OdDrDil03MCb1H2o9nJ1x4/5LYw7byDE0= github.com/sirupsen/logrus v1.9.0 h1:trlNQbNUG3OdDrDil03MCb1H2o9nJ1x4/5LYw7byDE0=
+48 -46
View File
@@ -7,9 +7,6 @@ import (
"errors" "errors"
"fmt" "fmt"
"io" "io"
"net/url"
"os"
"path/filepath"
"regexp" "regexp"
"strings" "strings"
@@ -22,22 +19,29 @@ type Defaults struct {
Provider string `json:"provider"` Provider string `json:"provider"`
Model string `json:"model"` Model string `json:"model"`
Thinking string `json:"thinking"` Thinking string `json:"thinking"`
LLMURL string `json:"llm_url"`
} }
// Configure validates and atomically stores only local non-secret Pi defaults. var internalIdentityHeaders = []string{
func Configure(ctx context.Context, runner Runner, path string, value Defaults) error { "-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) { if !choicePattern.MatchString(value.Provider) || !choicePattern.MatchString(value.Model) {
return errors.New("provider and model must be supported identifiers") return errors.New("provider and model must be supported identifiers")
} }
if value.Thinking != "low" && value.Thinking != "medium" && value.Thinking != "high" { if value.Thinking != "low" && value.Thinking != "medium" && value.Thinking != "high" {
return errors.New("thinking must be low, medium, or high") return errors.New("thinking must be low, medium, or high")
} }
u, err := url.Parse(value.LLMURL) before, err := renderedCore(ctx, runner)
if err != nil || (u.Scheme != "https" && u.Scheme != "http") || u.Host == "" || u.User != nil || u.RawQuery != "" || u.Fragment != "" { if err != nil { return err }
return errors.New("LLM endpoint must be an http(s) URL without credentials, query, or fragment") args := append([]string{"exec", "-T", "core", "curl", "-fsS"}, internalIdentityHeaders...)
} args = append(args, "http://127.0.0.1:8787/models")
models, err := runCompose(ctx, runner, "exec", "-T", "core", "curl", "-fsS", "http://127.0.0.1:8787/models") models, err := runCompose(ctx, runner, args...)
if err != nil { if err != nil {
return commandError("Pi options check", models, err) return commandError("Pi options check", models, err)
} }
@@ -59,39 +63,19 @@ func Configure(ctx context.Context, runner Runner, path string, value Defaults)
if !found { if !found {
return errors.New("provider/model is not in Pi options") return errors.New("provider/model is not in Pi options")
} }
return writeJSON(path, value) 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...)
func writeJSON(path string, value any) error { settingsArgs = append(settingsArgs, "http://127.0.0.1:8787/settings")
contents, err := json.MarshalIndent(value, "", " ") settings, err := runCompose(ctx, runner, settingsArgs...)
if err != nil { if err != nil { return commandError("Pi installation settings read-back", settings, err) }
return err var saved Defaults
} if json.Unmarshal([]byte(settings.Stdout), &saved) != nil || saved.Provider != value.Provider || saved.Model != value.Model || saved.Thinking != value.Thinking {
contents = append(contents, '\n') return errors.New("Pi installation settings read-back did not match requested provider, model, and thinking")
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")
} }
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 return nil
} }
@@ -131,7 +115,7 @@ func Doctor(ctx context.Context, runner Runner) error {
return commandError("Pi preflight check", result, err) 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. // 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 return err
} }
for _, path := range []string{"health", "models", "settings"} { 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 { if err != nil {
return commandError("Pi smoke check", result, err) 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 { if err := json.Unmarshal([]byte(result.Stdout), &payload); err != nil {
return fmt.Errorf("Pi smoke check returned invalid %s response", path) 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) 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 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) { func renderedCore(ctx context.Context, runner Runner) (Image, error) {
result, err := runCompose(ctx, runner, "config", "--format", "json") result, err := runCompose(ctx, runner, "config", "--format", "json")
if err != nil { if err != nil {
+7 -7
View File
@@ -2,7 +2,6 @@ package pi
import ( import (
"context" "context"
"path/filepath"
"strings" "strings"
"testing" "testing"
) )
@@ -17,16 +16,17 @@ func TestDoctorRequiresExternalEndpointAuthPiStateAndHealth(t *testing.T) {
} }
} }
func TestConfigureValidatesBackendModelOptionsAndWritesNoSecrets(t *testing.T) { func TestConfigureValidatesBackendModelOptionsWritesRealCoreSettingsAndUsesUpstreamIdentity(t *testing.T) {
fake := newFakeRunner() fake := newFakeRunner()
path := filepath.Join(t.TempDir(), "pi-defaults.json") if err := Configure(context.Background(), fake, Defaults{Provider: "provider", Model: "model", Thinking: "medium"}); err != nil {
if err := Configure(context.Background(), fake, path, Defaults{Provider: "provider", Model: "model", Thinking: "medium", LLMURL: "https://llm.example.invalid"}); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if got := string(readStateBytes(t, path)); strings.Contains(got, "secret") || !strings.Contains(got, "llm.example.invalid") { assertCalled(t, fake.calls, "node /app/backend/dist/settings/settings-cli.js --provider provider --model model --thinking medium")
t.Fatalf("defaults=%q", got) 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") t.Fatal("expected unknown model rejection")
} }
} }
@@ -0,0 +1,15 @@
//go:build !windows
package pi
import "os"
// durableReplace acknowledges both the data file and its directory entry. A successful return
// is the strongest atomic replacement guarantee supported by Unix filesystems.
func durableReplace(temporary, target, directory string) error {
if err := os.Rename(temporary, target); err != nil { return err }
dir, err := os.Open(directory)
if err != nil { return err }
defer dir.Close()
return dir.Sync()
}
@@ -0,0 +1,15 @@
//go:build windows
package pi
import "golang.org/x/sys/windows"
// MoveFileEx requests replacement and write-through on Windows. Directory fsync is not exposed
// by the Windows API in the same form as Unix, so callers must not claim a stronger guarantee.
func durableReplace(temporary, target, _ string) error {
from, err := windows.UTF16PtrFromString(temporary)
if err != nil { return err }
to, err := windows.UTF16PtrFromString(target)
if err != nil { return err }
return windows.MoveFileEx(from, to, windows.MOVEFILE_REPLACE_EXISTING|windows.MOVEFILE_WRITE_THROUGH)
}
@@ -0,0 +1,16 @@
//go:build !windows
package pi
import (
"errors"
"os"
"syscall"
)
func processAlive(pid int) bool {
process, err := os.FindProcess(pid)
if err != nil { return false }
err = process.Signal(syscall.Signal(0))
return err == nil || errors.Is(err, syscall.EPERM)
}
@@ -0,0 +1,14 @@
//go:build windows
package pi
import "golang.org/x/sys/windows"
func processAlive(pid int) bool {
handle, err := windows.OpenProcess(windows.PROCESS_QUERY_LIMITED_INFORMATION, false, uint32(pid))
if err != nil { return err == windows.ERROR_ACCESS_DENIED }
defer windows.CloseHandle(handle)
var code uint32
if windows.GetExitCodeProcess(handle, &code) != nil { return true }
return code == 259 // STILL_ACTIVE
}
+77 -40
View File
@@ -5,13 +5,15 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"crypto/sha256"
"os" "os"
"path/filepath" "path/filepath"
"runtime" "sort"
"strings"
"time" "time"
) )
const stateFileVersion = 1 const stateFileVersion = 2
// Phase describes the durable point reached by a Pi update. // Phase describes the durable point reached by a Pi update.
type Phase string type Phase string
@@ -32,6 +34,7 @@ type Image struct {
Reference string `json:"reference"` Reference string `json:"reference"`
Volumes []string `json:"volumes"` Volumes []string `json:"volumes"`
Mounts []Mount `json:"mounts"` Mounts []Mount `json:"mounts"`
MountFingerprint string `json:"mount_fingerprint"`
ConfigurationSHA string `json:"configuration_sha256,omitempty"` ConfigurationSHA string `json:"configuration_sha256,omitempty"`
} }
@@ -39,9 +42,10 @@ type Image struct {
type Mount struct { type Mount struct {
Type string `json:"type"` Type string `json:"type"`
Name string `json:"name,omitempty"` Name string `json:"name,omitempty"`
Source string `json:"source"` SourceSHA256 string `json:"source_sha256"`
Destination string `json:"destination"` Destination string `json:"destination"`
RW bool `json:"rw"` RW bool `json:"rw"`
Options string `json:"options,omitempty"`
} }
// Target records the immutable input selected by the operator. Source is either build or a // 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 { if err := json.Unmarshal(contents, &state); err != nil {
return State{}, errors.New("update recovery state is invalid") 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{}, errors.New("update recovery state is incomplete")
} }
return state, nil return state, nil
} }
func writeState(path string, state State) error { 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") return errors.New("refusing to write incomplete update recovery state")
} }
state.Version = stateFileVersion state.Version = stateFileVersion
@@ -89,45 +93,52 @@ func writeState(path string, state State) error {
return fmt.Errorf("encode update recovery state: %w", err) return fmt.Errorf("encode update recovery state: %w", err)
} }
contents = append(contents, '\n') contents = append(contents, '\n')
directory := filepath.Dir(path) if err := writeFileDurably(path, ".update-state-", contents); err != nil {
if err := os.MkdirAll(directory, 0o700); err != nil { return fmt.Errorf("could not durably write update recovery state: %w", err)
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()
}
} }
return nil 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 } type updateLock struct{ path string }
var ErrLockHeld = errors.New("another Pi update or rollback is already in progress")
func acquireLock(statePath string) (*updateLock, error) { func acquireLock(statePath string) (*updateLock, error) {
if err := os.MkdirAll(filepath.Dir(statePath), 0o700); err != nil { if err := os.MkdirAll(filepath.Dir(statePath), 0o700); err != nil {
return nil, errors.New("could not create Pi update recovery directory") return nil, errors.New("could not create Pi update recovery directory")
@@ -135,10 +146,36 @@ func acquireLock(statePath string) (*updateLock, error) {
path := statePath + ".lock" path := statePath + ".lock"
if err := os.Mkdir(path, 0o700); err != nil { if err := os.Mkdir(path, 0o700); err != nil {
if errors.Is(err, os.ErrExist) { 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") 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 return &updateLock{path: path}, nil
} }
func (l *updateLock) Release() { _ = os.Remove(l.path) } func (l *updateLock) Release() { _ = os.Remove(filepath.Join(l.path, "owner.json")); _ = os.Remove(l.path) }
// reclaimDeadLocalLock is deliberately conservative: a malformed, remote, or merely old lock
// is recovery-required. Only a process we can prove is gone on this machine is reclaimed.
func reclaimDeadLocalLock(path string) bool {
contents, err := os.ReadFile(filepath.Join(path, "owner.json"))
if err != nil { return false }
var owner lockOwner
if json.Unmarshal(contents, &owner) != nil || owner.PID <= 0 || owner.Host == "" { return false }
host, err := os.Hostname()
if err != nil || owner.Host != host { return false }
if processAlive(owner.PID) { return false }
if err := os.Remove(filepath.Join(path, "owner.json")); err != nil { return false }
return os.Remove(path) == nil
}
+84 -52
View File
@@ -10,14 +10,16 @@ import (
"sort" "sort"
"strings" "strings"
"time" "time"
"github.com/distribution/reference"
) )
var ( var (
ErrConfirmationRequired = errors.New("update requires --yes after reviewing the planned Pi version") 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") 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") 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.-]+)?$`) 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. // 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. // 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) lock, err := acquireLock(request.StatePath)
if err != nil { if err != nil {
return Result{StatePath: request.StatePath}, err 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 return Result{StatePath: request.StatePath}, ErrConfirmationRequired
} }
if !versionPattern.MatchString(request.Version) { if !versionPattern.MatchString(request.Version) {
return Result{StatePath: request.StatePath}, errors.New("Pi version must be an explicit pinned version") return Result{StatePath: request.StatePath}, fmt.Errorf("%w: Pi version must be an explicit pinned version", ErrInvalidRequest)
}
if request.Source == "" {
request.Source = BuildSource
} }
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 { 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) { if request.Source == PullSource {
return Result{StatePath: request.StatePath}, errors.New("pulled Pi image must use an immutable sha256 digest") 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 { if old, err := readState(request.StatePath); err == nil && old.Phase != PhaseVerified && old.Phase != PhaseRolledBack && old.Phase != PhaseNoop {
return Result{StatePath: request.StatePath}, ErrInterruptedUpdate return Result{StatePath: request.StatePath}, ErrInterruptedUpdate
} else if err != nil && !errors.Is(err, os.ErrNotExist) { } else if err != nil && !errors.Is(err, os.ErrNotExist) {
return Result{StatePath: request.StatePath}, err 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) running, err := activeSessions(ctx, runner)
if err != nil { if err != nil {
return Result{StatePath: request.StatePath}, err return Result{StatePath: request.StatePath}, err
} }
frontendStopped := false
if running { if running {
if !request.Drain { if !request.Drain {
return Result{StatePath: request.StatePath}, ErrActiveSessions 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++ { for attempts := 0; attempts < 30; attempts++ {
running, err = activeSessions(ctx, runner) running, err = activeSessions(ctx, runner)
if err != nil { if err != nil {
@@ -134,41 +133,34 @@ func Update(ctx context.Context, runner Runner, request Request) (Result, error)
} }
state.Phase = PhaseBuilding state.Phase = PhaseBuilding
if err := writeState(request.StatePath, state); err != nil { if err := writeState(request.StatePath, state); err != nil { return Result{Phase: PhaseFailed, StatePath: request.StatePath}, err }
return Result{StatePath: request.StatePath}, err
}
if err := prepareCandidate(ctx, runner, request, previous.Reference); err != nil { if err := prepareCandidate(ctx, runner, request, previous.Reference); err != nil {
state.Phase, state.Error = PhaseFailed, "candidate image preparation failed" return compensate(ctx, runner, request.StatePath, state, err)
_ = writeState(request.StatePath, state)
return Result{Phase: PhaseFailed, StatePath: request.StatePath}, err
} }
if frontendStopped { if request.Drain {
running, err = activeSessions(ctx, runner) running, err = activeSessions(ctx, runner)
if err != nil { if err != nil {
return Result{Phase: PhaseFailed, StatePath: request.StatePath}, err return compensate(ctx, runner, request.StatePath, state, err)
} }
if running { if running {
state.Phase, state.Error = PhaseFailed, "new session admitted while draining" return compensate(ctx, runner, request.StatePath, state, ErrActiveSessions)
_ = writeState(request.StatePath, state)
return Result{Phase: PhaseFailed, StatePath: request.StatePath}, ErrActiveSessions
} }
} }
if err := recreateCore(ctx, runner); err != nil { if err := recreateCore(ctx, runner); err != nil {
state.Phase, state.Error = PhaseFailed, "core recreation failed" return compensate(ctx, runner, request.StatePath, state, err)
_ = writeState(request.StatePath, state)
return Result{Phase: PhaseFailed, StatePath: request.StatePath}, err
} }
state.Phase = PhaseRecreated 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 { 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 { 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, "" state.Phase, state.Error = PhaseVerified, ""
if err := writeState(request.StatePath, state); err != nil { 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 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 { if err := restore(ctx, runner, state.Previous); err != nil {
state.Phase, state.Error = PhaseFailed, "rollback failed" 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 return Result{Phase: PhaseFailed, StatePath: statePath}, err
} }
state.Phase, state.Error = PhaseRolledBack, "" state.Phase, state.Error = PhaseRolledBack, ""
if err := writeState(statePath, state); err != nil { 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 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 { if restoreErr := restore(ctx, runner, state.Previous); restoreErr != nil {
state.Phase, state.Error = PhaseFailed, "candidate verification and automatic rollback failed" state.Phase, state.Error = PhaseFailed, "candidate verification and automatic rollback failed"
_ = writeState(statePath, state) if writeErr := writeState(statePath, state); writeErr != nil {
return Result{Phase: PhaseFailed, StatePath: statePath}, fmt.Errorf("candidate verification failed; automatic rollback also failed") 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, "" state.Phase, state.Error = PhaseRolledBack, ""
_ = writeState(statePath, state) if writeErr := writeState(statePath, state); writeErr != nil {
return Result{Phase: PhaseRolledBack, StatePath: statePath}, fmt.Errorf("candidate verification failed; previous core image was restored") 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 { func sourceValue(request Request) string {
@@ -217,6 +219,29 @@ func sourceValue(request Request) string {
return string(BuildSource) 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) { func activeSessions(ctx context.Context, runner Runner) (bool, error) {
result, err := runCompose(ctx, runner, "exec", "-T", "core", "tht", "session", "list", "--json") result, err := runCompose(ctx, runner, "exec", "-T", "core", "tht", "session", "list", "--json")
if err != nil { if err != nil {
@@ -251,20 +276,27 @@ func runningImage(ctx context.Context, runner Runner, reference string) (Image,
if err != nil { if err != nil {
return Image{}, commandError("core volume check", mounts, err) return Image{}, commandError("core volume check", mounts, err)
} }
var contract []Mount var raw []struct {
if err := json.Unmarshal([]byte(mounts.Stdout), &contract); err != nil { 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") 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") return Image{}, errors.New("core has no persistence mounts to preserve")
} }
volumes := make([]string, 0, len(contract)) contract := make([]Mount, 0, len(raw))
for _, mount := range contract { 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 != "" { if mount.Type == "volume" && mount.Name != "" {
volumes = append(volumes, 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 { func prepareCandidate(ctx context.Context, runner Runner, request Request, reference string) error {
@@ -382,7 +414,7 @@ func sameMounts(left, right []Mount) bool {
return false return false
} }
key := func(m Mount) string { 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)) a, b := make([]string, len(left)), make([]string, len(right))
for i := range left { for i := range left {
+22 -8
View File
@@ -3,6 +3,7 @@ package pi
import ( import (
"context" "context"
"errors" "errors"
"fmt"
"io" "io"
"os" "os"
"path/filepath" "path/filepath"
@@ -103,10 +104,9 @@ func TestUpdateDoesNotRecreateWhenPreflightOrBuildFails(t *testing.T) {
if err == nil { if err == nil {
t.Fatal("Update() error = nil, want failure") t.Fatal("Update() error = nil, want failure")
} }
if result.Phase == PhaseRolledBack { if failure == "preflight" && result.Phase == PhaseRolledBack { t.Fatalf("preflight failure unexpectedly rolled back: %+v", result) }
t.Fatalf("pre-recreate 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") }
assertNotCalled(t, fake.calls, "force-recreate")
}) })
} }
} }
@@ -162,7 +162,7 @@ func TestRollbackRestoresInterruptedOrPreviouslyRecordedState(t *testing.T) {
func TestUpdateRefusesToOverwriteInterruptedRecoveryState(t *testing.T) { func TestUpdateRefusesToOverwriteInterruptedRecoveryState(t *testing.T) {
fake := newFakeRunner() fake := newFakeRunner()
statePath := filepath.Join(t.TempDir(), "state.json") 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}) _, err := Update(context.Background(), fake, Request{StatePath: statePath, Version: "0.81.0", Source: BuildSource, Confirm: true})
if !errors.Is(err, ErrInterruptedUpdate) { if !errors.Is(err, ErrInterruptedUpdate) {
t.Fatalf("Update() error = %v, want interrupted update error", err) 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" { if len(image.Mounts) != 3 || image.Mounts[0].Type != "bind" || image.Mounts[0].Destination != "/data" {
t.Fatalf("mounts = %#v", image.Mounts) 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 { type fakeRunner struct {
@@ -188,12 +199,13 @@ type fakeRunner struct {
version string version string
activeSessions bool activeSessions bool
built bool built bool
currentImage string
volumes []string volumes []string
mountsJSON string mountsJSON string
} }
func newFakeRunner() *fakeRunner { 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) { 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) f.calls = append(f.calls, call)
if strings.Contains(call, "image tag sha256:old") { if strings.Contains(call, "image tag sha256:old") {
f.fail = "" f.fail = ""
f.currentImage = "sha256:old"
} }
if f.fail == "preflight" && strings.Contains(call, "config --format json") { if f.fail == "preflight" && strings.Contains(call, "config --format json") {
return compose.Result{ExitCode: 1}, errors.New("provider token=secret") 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") { 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") 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") return compose.Result{ExitCode: 1}, errors.New("smoke token=secret")
} }
switch { switch {
@@ -223,7 +236,7 @@ func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose
case strings.Contains(call, "ps -q core"): case strings.Contains(call, "ps -q core"):
return compose.Result{Stdout: "core-container\n"}, nil return compose.Result{Stdout: "core-container\n"}, nil
case strings.Contains(call, "inspect --format {{.Image}}"): 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}}"): case strings.Contains(call, "inspect --format {{json .Mounts}}"):
if f.mountsJSON != "" { if f.mountsJSON != "" {
return compose.Result{Stdout: f.mountsJSON}, nil 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"): case strings.Contains(call, "compose build"):
f.built = true f.built = true
f.version = "0.81.0" f.version = "0.81.0"
f.currentImage = "sha256:candidate"
return compose.Result{}, nil return compose.Result{}, nil
case strings.Contains(call, "pi --version"): case strings.Contains(call, "pi --version"):
return compose.Result{Stdout: f.version + "\n"}, nil return compose.Result{Stdout: f.version + "\n"}, nil