diff --git a/backend/src/auth/diagnostic-command.ts b/backend/src/auth/diagnostic-command.ts index c5883a6d..75bd1fee 100644 --- a/backend/src/auth/diagnostic-command.ts +++ b/backend/src/auth/diagnostic-command.ts @@ -1,31 +1,41 @@ import { fileURLToPath } from "node:url"; import { resolve } from "node:path"; import { loadConfig, type AppConfig } from "../config.js"; -import { secretValue } from "../config/secret-bundle.js"; +import { loadSecretBundle, SECRET_BUNDLE_KEYS, secretValue } from "../config/secret-bundle.js"; import { createAuthentikGroupCatalog } from "./authentik-group-catalog.js"; import { createCurrentLocalUserRegistryResolver, type LocalUserRegistry } from "./local-registry.js"; import { createOidcProtocol, OidcDeviceFlowUnavailableError, type OidcProtocol } from "./oidc-client.js"; import { createAuthDiagnoser, type AuthDiagnoser, type AuthDiagnostic, type AuthDiagnostics } from "./diagnostics.js"; +import { decodeAuthDiagnostics, type GroupCatalog } from "./group-catalog.js"; import type { LoadedAuthConfig } from "./types.js"; const AUTH_SECRET_REFERENCES = ["THT_OIDC_CLIENT_SECRET", "THT_AUTHENTIK_API_TOKEN"] as const; -function configuredSecretValues(config: AppConfig): readonly string[] { - const values: string[] = []; - for (const reference of AUTH_SECRET_REFERENCES) { +export function configuredSecretValues(config: AppConfig): readonly string[] { + const values = new Set(); + if (config.secretsFile) { try { - const value = secretValue(config, reference); - if (value !== undefined) values.push(value); + for (const value of loadSecretBundle(config.secretsFile).values()) values.add(value); + } catch { + return []; + } + } + for (const reference of SECRET_BUNDLE_KEYS) { + try { + const value = secretValue({ secretFiles: config.secretFiles }, reference); + if (value !== undefined) values.add(value); } catch { // The fixed report below is the only externally-visible failure surface. } } - return values; + return [...values]; } export interface ConfiguredAuthDiagnoserOptions { localUserRegistry?: (loaded: LoadedAuthConfig) => LocalUserRegistry | undefined; oidcProtocol?: (loaded: LoadedAuthConfig) => OidcProtocol | undefined; + groupCatalog?: (loaded: LoadedAuthConfig) => GroupCatalog | undefined; + sessionRootValidator?: (root: string) => void | Promise; } /** Builds the one shared auth diagnostic implementation used by app routes and the one-shot CLI. */ @@ -70,7 +80,7 @@ export function createConfiguredAuthDiagnoser( } catch { return undefined; } })() : undefined; - const groupCatalog = current?.value.mode === "oidc" ? (() => { + const groupCatalog = current?.value.mode === "oidc" ? options.groupCatalog?.(current) ?? (() => { try { const token = secretValues().get("THT_AUTHENTIK_API_TOKEN"); return token === undefined ? undefined : createAuthentikGroupCatalog({ @@ -82,6 +92,7 @@ export function createConfiguredAuthDiagnoser( const report = await createAuthDiagnoser({ authMode: config.authMode, authStateRoot: config.authStateRoot, + ...(options.sessionRootValidator === undefined ? {} : { sessionRootValidator: options.sessionRootValidator }), authentication: config.authentication, secrets: secretValues(), localUserRegistry: current?.value.mode === "local" @@ -108,10 +119,21 @@ export function createConfiguredAuthDiagnoser( ); // Exact names only: unrelated provider groups are neither emitted nor retained. const mappedRoles = new Set(); - for (const group of identity.groups) { - for (const role of current.value.authorization.groupRoles[group] ?? []) mappedRoles.add(role); + for (const [configuredGroup, roles] of Object.entries(current.value.authorization.groupRoles)) { + if (!identity.groups.includes(configuredGroup)) continue; + for (const role of roles) mappedRoles.add(role); + } + if (mappedRoles.size === 0) { + return { + ready: false, + mode: "oidc", + checks: [{ + level: "error", + code: "oidc_groups_claim_invalid", + message: "The OIDC device-flow identity could not be validated.", + }], + }; } - void mappedRoles; return report; } catch (error) { return { @@ -188,18 +210,22 @@ export async function runDiagnosticCommand( return 2; } let report: AuthDiagnostics; + const secrets = dependencies.secretValues ?? []; try { report = await dependencies.diagnoser.inspect({ live: true, ...(options.interactive ? { interactive: true, - presentDeviceCode: (uri: string, code: string) => dependencies.stderr(`Open ${uri} and enter code ${code}`), + presentDeviceCode: (uri: string, code: string) => dependencies.stderr( + redact(`Open ${uri} and enter code ${code}`, secrets), + ), } : {}), }); } catch { report = genericFailure(); } - const safe = redactedReport(report, dependencies.secretValues ?? []); + const decoded = decodeAuthDiagnostics(report) ?? genericFailure(); + const safe = decodeAuthDiagnostics(redactedReport(decoded, secrets)) ?? genericFailure(); dependencies.stdout(`${JSON.stringify(safe)}\n`); return safe.ready ? 0 : 1; } diff --git a/backend/src/auth/group-catalog.ts b/backend/src/auth/group-catalog.ts index 281bf66d..eb06f7ca 100644 --- a/backend/src/auth/group-catalog.ts +++ b/backend/src/auth/group-catalog.ts @@ -30,6 +30,66 @@ export interface AuthDiagnostics { checks: readonly AuthDiagnostic[]; } +const diagnosticCodes = new Set([ + "auth_ready", "auth_config_incomplete", "auth_config_invalid", "auth_session_store_invalid", + "local_user_registry_invalid", "local_admin_missing", "oidc_secret_missing", + "oidc_discovery_unreachable", "oidc_issuer_mismatch", "oidc_jwks_unreachable", + "oidc_group_catalog_unreachable", "oidc_group_catalog_unauthorized", "oidc_mapped_group_missing", + "oidc_mapped_group_ambiguous", "oidc_groups_claim_invalid", "oidc_device_flow_unavailable", +]); +const diagnosticModes = new Set(["local", "oidc", "upstream", "none", "mock"]); +const fieldCodes = new Set(["oidc_mapped_group_missing", "oidc_mapped_group_ambiguous"]); + +function exactObject(value: unknown, keys: readonly string[]): Record | undefined { + if (!value || typeof value !== "object" || Array.isArray(value)) return undefined; + const source = value as Record; + const actual = Object.keys(source); + return actual.length === keys.length && actual.every((key) => keys.includes(key)) ? source : undefined; +} + +function safeText(value: unknown): value is string { + return typeof value === "string" && value.length > 0 && value.length <= 512 + && value.trim() === value && !/\p{Cc}/u.test(value); +} + +/** Strict decoder for the machine contract shared with tht and the frontend. */ +export function decodeAuthDiagnostics(value: unknown): AuthDiagnostics | undefined { + const source = exactObject(value, ["ready", "mode", "checks"]); + if (!source || typeof source.ready !== "boolean" || typeof source.mode !== "string" + || !diagnosticModes.has(source.mode as AuthDiagnostics["mode"]) + || !Array.isArray(source.checks) || source.checks.length === 0 || source.checks.length > 129) return undefined; + const seen = new Set(); + const checks: AuthDiagnostic[] = []; + for (const value of source.checks) { + const raw = value && typeof value === "object" && !Array.isArray(value) + ? value as Record + : undefined; + const check = raw && exactObject(raw, raw.field === undefined + ? ["level", "code", "message"] + : ["level", "code", "message", "field"]); + if (!check || (check.level !== "error" && check.level !== "info") + || typeof check.code !== "string" || !diagnosticCodes.has(check.code as AuthDiagnosticCode) + || !safeText(check.message) || (check.field !== undefined && !safeText(check.field))) return undefined; + const code = check.code as AuthDiagnosticCode; + if (check.field !== undefined && !fieldCodes.has(code)) return undefined; + const key = `${code}\u0000${check.field ?? ""}`; + if (seen.has(key)) return undefined; + seen.add(key); + checks.push({ + level: check.level, + code, + message: check.message, + ...(check.field === undefined ? {} : { field: check.field }), + }); + } + if (source.ready) { + if (checks.length !== 1 || checks[0].level !== "info" || checks[0].code !== "auth_ready" + || checks[0].field !== undefined) return undefined; + } else if (!checks.some(({ level }) => level === "error") + || checks.some(({ code }) => code === "auth_ready")) return undefined; + return { ready: source.ready, mode: source.mode as AuthDiagnostics["mode"], checks }; +} + /** A provider-specific proof that only the configured authorization groups exist. */ export interface GroupCatalog { verifyConfiguredGroups(names: readonly string[], signal: AbortSignal): Promise; diff --git a/backend/src/auth/oidc-client.ts b/backend/src/auth/oidc-client.ts index a4a60709..70c8a833 100644 --- a/backend/src/auth/oidc-client.ts +++ b/backend/src/auth/oidc-client.ts @@ -611,23 +611,37 @@ export function createOidcProtocol(options: OidcProtocolOptions): OidcProtocol { signal.throwIfAborted(); }, async verifyDeviceFlow(signal, present) { + const maximumDeadline = AbortSignal.timeout(MAX_DEVICE_FLOW_TIMEOUT_MS); + const operationSignal = AbortSignal.any([signal, maximumDeadline]); let config: Configuration; try { + operationSignal.throwIfAborted(); config = await configuration(); + operationSignal.throwIfAborted(); const endpoint = config.serverMetadata().device_authorization_endpoint; httpsEndpoint(endpoint); } catch (error) { if (error instanceof OidcIssuerMismatchError || error instanceof OidcProviderUnavailableError) throw error; throw new OidcDeviceFlowUnavailableError(); } - const deadline = AbortSignal.timeout(MAX_DEVICE_FLOW_TIMEOUT_MS); - const deviceSignal = AbortSignal.any([signal, deadline]); try { - deviceSignal.throwIfAborted(); - const device = await initiateDeviceAuthorization(config, { scope: options.scopes.join(" ") }); + operationSignal.throwIfAborted(); + const device = await awaitWithAbort( + initiateDeviceAuthorization(config, { scope: options.scopes.join(" ") }), + operationSignal, + ); if (!text(device.verification_uri, 2048) || !text(device.user_code, 256)) { throw new OidcDeviceFlowUnavailableError(); } + const providerLifetimeMs = device.expires_in * 1000; + if (!Number.isSafeInteger(providerLifetimeMs) || providerLifetimeMs <= 0) { + throw new OidcDeviceFlowUnavailableError(); + } + const deviceSignal = AbortSignal.any([ + signal, + maximumDeadline, + AbortSignal.timeout(Math.min(providerLifetimeMs, MAX_DEVICE_FLOW_TIMEOUT_MS)), + ]); const verificationUri = httpsEndpoint(device.verification_uri); present(verificationUri.href, device.user_code); const tokens = await pollDeviceAuthorizationGrant(config, device, undefined, { signal: deviceSignal }); diff --git a/backend/src/routes/workspaces.ts b/backend/src/routes/workspaces.ts index f6b04108..a2b4a563 100644 --- a/backend/src/routes/workspaces.ts +++ b/backend/src/routes/workspaces.ts @@ -19,6 +19,7 @@ import type { RuntimeBindings } from "../workspaces/runtime-renderer.js"; import type { ConnectorDiagnostics } from "../workspaces/diagnostics.js"; import { isPrincipalContext, requirePermission } from "../auth/authorization.js"; import type { AuthDiagnoser } from "../auth/diagnostics.js"; +import { decodeAuthDiagnostics, type AuthDiagnostics } from "../auth/group-catalog.js"; export type WorkspaceDiagnoser = ( workspace: WorkspaceDescriptor, @@ -55,6 +56,12 @@ const SAFE_MESSAGES = { semantic_index_incompatible: "Semantic index is incompatible with this workspace.", } as const; +function authenticationReport(value: unknown): AuthDiagnostics { + const report = decodeAuthDiagnostics(value); + if (!report) throw new Error("invalid authentication diagnostic report"); + return report; +} + function workspaceErrorCode(error: unknown): keyof typeof SAFE_MESSAGES { return error instanceof WorkspaceRegistryError ? error.code : "workspace_invalid"; } @@ -149,7 +156,7 @@ export function workspaceRoutes(app: FastifyInstance, deps: WorkspaceRoutesDeps) try { const { workspace } = workspacePayload.parse(request.body); const canonical = validateWorkspaceDescriptor(workspace); - const authentication = await deps.authDiagnoser.inspect({ live: false }); + const authentication = authenticationReport(await deps.authDiagnoser.inspect({ live: false })); return { workspace: canonical, contract: buildInstallationContract(canonical), @@ -231,10 +238,11 @@ export function workspaceRoutes(app: FastifyInstance, deps: WorkspaceRoutesDeps) deps.secretStore, ); try { - const [workspaceDiagnostics, authentication] = await Promise.all([ + const [workspaceDiagnostics, inspectedAuthentication] = await Promise.all([ deps.diagnose(operational, lease.bindings, { writeProbe: false }), deps.authDiagnoser.inspect({ live: true }), ]); + const authentication = authenticationReport(inspectedAuthentication); return { ...workspaceDiagnostics, activatable: workspaceDiagnostics.activatable && authentication.ready, diff --git a/backend/test/auth-diagnostic-command.test.ts b/backend/test/auth-diagnostic-command.test.ts index a049b97a..46a1cede 100644 --- a/backend/test/auth-diagnostic-command.test.ts +++ b/backend/test/auth-diagnostic-command.test.ts @@ -1,6 +1,69 @@ -import { expect, test, vi } from "vitest"; +import { chmodSync, mkdtempSync, rmSync, writeFileSync } from "node:fs"; +import { tmpdir } from "node:os"; +import { join } from "node:path"; +import { afterEach, expect, test, vi } from "vitest"; import type { AuthDiagnoser, AuthDiagnostics } from "../src/auth/diagnostics.js"; -import { runDiagnosticCommand } from "../src/auth/diagnostic-command.js"; +import { + configuredSecretValues, + createConfiguredAuthDiagnoser, + runDiagnosticCommand, +} from "../src/auth/diagnostic-command.js"; +import type { AppConfig } from "../src/config.js"; +import type { LoadedAuthConfig } from "../src/auth/types.js"; + +const roots: string[] = []; +afterEach(() => { for (const root of roots.splice(0)) rmSync(root, { recursive: true, force: true }); }); + +function secretBundle(contents: string): string { + const root = mkdtempSync(join(tmpdir(), "thothii-diagnostic-command-")); + roots.push(root); + const file = join(root, "secrets"); + writeFileSync(file, contents, { mode: 0o600 }); + chmodSync(file, 0o600); + return file; +} + +function configuredOidc(groups: readonly string[]): AuthDiagnoser { + const loaded: LoadedAuthConfig = { + revision: "a".repeat(64), + sourcePath: "/safe/auth.yaml", + value: { + version: 1, + mode: "oidc", + publicUrl: "https://thothii.example.test", + session: { regularTtlSeconds: 1, regularIdleSeconds: 1, rememberTtlSeconds: 1, rememberIdleSeconds: 1, oidcTtlSeconds: 1 }, + oidc: { + issuer: "https://issuer.example.test", clientId: "thothii", + clientSecretRef: "THT_OIDC_CLIENT_SECRET", scopes: ["openid"], groupsClaim: "groups", + }, + groupCatalog: { + driver: "authentik", baseUrl: "https://authentik.example.test", + apiTokenRef: "THT_AUTHENTIK_API_TOKEN", + }, + authorization: { groupRoles: { "Thoth Users": ["user"], "Thoth Administrators": ["admin"] } }, + }, + }; + const config = { + authMode: "oidc", + authStateRoot: "/safe/auth-state", + authentication: { current: () => loaded }, + secretsFile: secretBundle("THT_OIDC_CLIENT_SECRET=oidc-secret\nTHT_AUTHENTIK_API_TOKEN=catalog-secret\n"), + secretFiles: {}, + } as AppConfig; + return createConfiguredAuthDiagnoser(config, { + sessionRootValidator: async () => undefined, + groupCatalog: () => ({ verifyConfiguredGroups: async () => [] }), + oidcProtocol: () => ({ + diagnose: async () => undefined, + authorizationUrl: async () => new URL("https://issuer.example.test/authorize"), + callback: async () => { throw new Error("not used"); }, + verifyDeviceFlow: async () => ({ + issuer: "https://issuer.example.test", subject: "subject", groups, + tokenExpiresAt: new Date(Date.now() + 60_000), + }), + }), + } as any); +} const failure: AuthDiagnostics = { ready: false, @@ -65,3 +128,110 @@ test("delegates an interactive diagnostic to OIDC and keeps its device prompt on expect(stderr).toEqual(["Open https://issuer.example.test/device and enter code ABCD-EFGH"]); expect(JSON.parse(stdout.join(""))).toEqual(ready); }); + +test("fails a direct-group identity that maps to zero Thoth roles without exposing provider groups", async () => { + const unmapped = "Unmapped Provider Group SECRET-SENTINEL"; + const report = await configuredOidc([unmapped]).inspect({ + live: true, + interactive: true, + presentDeviceCode: () => undefined, + }); + + expect(report).toEqual({ + ready: false, + mode: "oidc", + checks: [{ + level: "error", + code: "oidc_groups_claim_invalid", + message: "The OIDC device-flow identity could not be validated.", + }], + }); + expect(JSON.stringify(report)).not.toContain(unmapped); +}); + +test("accepts a direct-group identity with at least one exact configured role mapping", async () => { + const report = await configuredOidc(["Thoth Users", "Unmapped Provider Group"]).inspect({ + live: true, + interactive: true, + presentDeviceCode: () => undefined, + }); + + expect(report).toMatchObject({ ready: true, checks: [{ code: "auth_ready" }] }); + expect(JSON.stringify(report)).not.toContain("Unmapped Provider Group"); +}); + +test("redacts every parsed mounted secret before writing an interactive prompt", async () => { + const modelSecret = "model-secret-SENTINEL"; + const deviceCode = "DEVICE-CODE-SENTINEL"; + const legacyVectorSecret = "legacy-vector-secret-SENTINEL"; + const config = { + secretsFile: secretBundle([ + `THT_MODEL_API_KEY=${modelSecret}`, + `THT_DWH_API_KEY=${deviceCode}`, + "THT_OIDC_CLIENT_SECRET=oidc-secret", + "THT_AUTHENTIK_API_TOKEN=catalog-secret", + "", + ].join("\n")), + secretFiles: { THT_VEC_API_KEY_SECRET_FILE: secretBundle(legacyVectorSecret) }, + } as AppConfig; + const secrets = configuredSecretValues(config); + const stdout: string[] = []; + const stderr: string[] = []; + + await runDiagnosticCommand(["--json", "--interactive"], { + diagnoser: { + inspect: async (request) => { + request.presentDeviceCode?.( + `https://issuer.example.test/device/${modelSecret}?legacy=${legacyVectorSecret}`, + deviceCode, + ); + return { ready: true, mode: "oidc", checks: [{ level: "info", code: "auth_ready", message: "Authentication is ready." }] }; + }, + }, + secretValues: secrets, + stdout: (line) => stdout.push(line), + stderr: (line) => stderr.push(line), + }); + + expect(secrets).toEqual(expect.arrayContaining([ + modelSecret, deviceCode, legacyVectorSecret, "oidc-secret", "catalog-secret", + ])); + expect(stdout.join("") + stderr.join("")).not.toContain(modelSecret); + expect(stdout.join("") + stderr.join("")).not.toContain(deviceCode); + expect(stdout.join("") + stderr.join("")).not.toContain(legacyVectorSecret); +}); + +test.each([ + { + ready: true, mode: "oidc", checks: [{ level: "error", code: "oidc_secret_missing", message: "attacker-message" }], + }, + { + ready: false, mode: "oidc", checks: [{ level: "info", code: "auth_ready", message: "attacker-message" }], + }, + { + ready: false, mode: "oidc", checks: [ + { level: "error", code: "oidc_secret_missing", message: "one" }, + { level: "error", code: "oidc_secret_missing", message: "duplicate-attacker" }, + ], + }, + { + ready: false, mode: "oidc", checks: [{ + level: "error", code: "oidc_secret_missing", message: "failure", field: "attacker-field-SENTINEL", + }], + }, +])("fails closed on a hostile semantic report %#", async (hostile) => { + const stdout: string[] = []; + const exitCode = await runDiagnosticCommand(["--json"], { + diagnoser: { inspect: async () => hostile as AuthDiagnostics }, + stdout: (line) => stdout.push(line), + stderr: () => undefined, + }); + + expect(exitCode).toBe(1); + expect(JSON.parse(stdout.join(""))).toEqual({ + ready: false, + mode: "none", + checks: [{ level: "error", code: "auth_config_invalid", message: "Authentication configuration is unavailable." }], + }); + expect(stdout.join("")).not.toMatch(/attacker|duplicate/i); +}); diff --git a/backend/test/oidc-client.test.ts b/backend/test/oidc-client.test.ts index 1051c691..c3da421a 100644 --- a/backend/test/oidc-client.test.ts +++ b/backend/test/oidc-client.test.ts @@ -1,5 +1,5 @@ import { createSign, generateKeyPairSync } from "node:crypto"; -import { expect, test } from "vitest"; +import { expect, test, vi } from "vitest"; import { createOidcProtocol, OidcIssuerMismatchError, @@ -170,6 +170,51 @@ test("refuses device flow when discovery has no safe device authorization endpoi .rejects.toBeInstanceOf(OidcDeviceFlowUnavailableError); }); +test("stops authorization-pending polling at the provider device expiry", async () => { + const caller = new AbortController(); + const timeout = vi.spyOn(AbortSignal, "timeout"); + try { + const tokenResponse = vi.fn(async () => Response.json( + { error: "authorization_pending", error_description: "pending" }, + { status: 400 }, + )); + const subject = protocol({ + discoveryMetadata: { device_authorization_endpoint: `${issuer}/device` }, + deviceResponse: async () => Response.json({ + device_code: "ephemeral-device-code", + user_code: "ABCD-EFGH", + verification_uri: `${issuer}/device`, + expires_in: 0.05, + interval: 0.01, + }), + tokenResponse, + }); + const completion = subject.verifyDeviceFlow!(caller.signal, () => undefined); + + expect(await settlesWithin(completion, 250)).toBe("rejected"); + expect(tokenResponse.mock.calls.length).toBeLessThanOrEqual(5); + expect(timeout).toHaveBeenCalledWith(10 * 60_000); + expect(timeout).toHaveBeenCalledWith(50); + await expect(completion).rejects.toThrow(OidcProtocolError); + } finally { + caller.abort(); + timeout.mockRestore(); + } +}); + +test("combines the caller cancellation with device-flow deadlines before provider requests", async () => { + const seen: URL[] = []; + const caller = new AbortController(); + caller.abort(new DOMException("caller deadline", "AbortError")); + const subject = protocol({ + seen, + discoveryMetadata: { device_authorization_endpoint: `${issuer}/device` }, + }); + + await expect(subject.verifyDeviceFlow!(caller.signal, () => undefined)).rejects.toThrow(OidcProtocolError); + expect(seen).toEqual([]); +}); + test("rejects a hanging discovery request at the provider transport deadline", async () => { let aborted = false; const subject = protocol({ diff --git a/backend/test/routes-workspaces.test.ts b/backend/test/routes-workspaces.test.ts index d32565a1..92368e82 100644 --- a/backend/test/routes-workspaces.test.ts +++ b/backend/test/routes-workspaces.test.ts @@ -250,6 +250,26 @@ test("aggregates one static and one live authentication report without reorderin expect(authDiagnoser.inspect).toHaveBeenNthCalledWith(2, { live: true }); }); +test("fails closed without reflecting a hostile authentication report", async () => { + const attacker = "attacker-field-SENTINEL"; + const authDiagnoser: AuthDiagnoser = { inspect: vi.fn(async () => ({ + ready: false, + mode: "oidc", + checks: [{ + level: "error", code: "oidc_secret_missing", message: "failure", field: attacker, + }], + } as AuthDiagnostics)) }; + const app = appFor(registryFake(), undefined, testSecretStore(), {}, authDiagnoser); + + const response = await app.inject({ + method: "POST", url: "/workspaces/validate", payload: { workspace }, + }); + + expect(response.statusCode).toBe(400); + expect(response.json()).toEqual({ code: "workspace_invalid", message: "Workspace request is invalid." }); + expect(response.body).not.toContain(attacker); +}); + test.each([1, 2])("rejects schema v%s at the validation boundary with a sanitized error", async (version) => { const legacy = { ...workspace, diff --git a/frontend/src/api/workspaces.test.ts b/frontend/src/api/workspaces.test.ts index e6f96812..f783381d 100644 --- a/frontend/src/api/workspaces.test.ts +++ b/frontend/src/api/workspaces.test.ts @@ -136,3 +136,26 @@ test("decodes the shared authentication diagnostics on static validation and liv authentication, }); }); + +test.each([ + ["ready with error", { ready: true, mode: "oidc", checks: [{ level: "error", code: "oidc_secret_missing", message: "failure" }] }], + ["failed with ready", { ready: false, mode: "oidc", checks: [{ level: "info", code: "auth_ready", message: "ready" }] }], + ["failed without error", { ready: false, mode: "oidc", checks: [{ level: "info", code: "auth_config_invalid", message: "info" }] }], + ["duplicate", { ready: false, mode: "oidc", checks: [ + { level: "error", code: "oidc_secret_missing", message: "one" }, + { level: "error", code: "oidc_secret_missing", message: "two" }, + ] }], + ["attacker field", { ready: false, mode: "oidc", checks: [{ level: "error", code: "oidc_secret_missing", message: "failure", field: "attacker-field-SENTINEL" }] }], + ["control", { ready: false, mode: "oidc", checks: [{ level: "error", code: "oidc_mapped_group_missing", message: "failure", field: "bad\u0085field" }] }], + ["unexpected property", { ready: false, mode: "oidc", checks: [{ level: "error", code: "oidc_secret_missing", message: "failure", attacker: "field" }] }], +])("rejects hostile authentication diagnostics: %s", async (_name, authentication) => { + server.use(http.post("/api/workspaces/validate", () => HttpResponse.json({ + workspace, + contract: {}, + activatable: false, + diagnostics: [], + authentication, + }))); + + await expect(validateWorkspace(workspace)).rejects.toThrow("invalid authentication diagnostics"); +}); diff --git a/frontend/src/api/workspaces.ts b/frontend/src/api/workspaces.ts index e8c0b382..f104a89b 100644 --- a/frontend/src/api/workspaces.ts +++ b/frontend/src/api/workspaces.ts @@ -225,21 +225,36 @@ const authDiagnosticCodes = new Set([ ]); function text(value: unknown): value is string { - return typeof value === "string" && value.length > 0 && value.length <= 512 && !/[\u0000-\u001f\u007f]/.test(value); + return typeof value === "string" && value.length > 0 && value.length <= 512 + && value.trim() === value && !/\p{Cc}/u.test(value); } function decodeAuthentication(value: unknown): AuthDiagnostics { - const source = object(value); + const source = exactObject(value, ["ready", "mode", "checks"]); if (!source || typeof source.ready !== "boolean" - || !["local", "oidc", "upstream", "none", "mock"].includes(String(source.mode)) - || !Array.isArray(source.checks)) throw new Error("Workspace API returned invalid authentication diagnostics"); + || typeof source.mode !== "string" + || !["local", "oidc", "upstream", "none", "mock"].includes(source.mode) + || !Array.isArray(source.checks) || source.checks.length === 0 || source.checks.length > 129) { + throw new Error("Workspace API returned invalid authentication diagnostics"); + } + const seen = new Set(); const checks = source.checks.map((item): AuthDiagnostic => { - const check = object(item); + const raw = object(item); + const check = raw && exactObject(raw, raw.field === undefined + ? ["level", "code", "message"] + : ["level", "code", "message", "field"]); if (!check || (check.level !== "error" && check.level !== "info") || typeof check.code !== "string" || !authDiagnosticCodes.has(check.code as AuthDiagnosticCode) || !text(check.message) || (check.field !== undefined && !text(check.field))) { throw new Error("Workspace API returned invalid authentication diagnostics"); } + if (check.field !== undefined + && check.code !== "oidc_mapped_group_missing" && check.code !== "oidc_mapped_group_ambiguous") { + throw new Error("Workspace API returned invalid authentication diagnostics"); + } + const key = `${check.code}\u0000${check.field ?? ""}`; + if (seen.has(key)) throw new Error("Workspace API returned invalid authentication diagnostics"); + seen.add(key); return { level: check.level, code: check.code as AuthDiagnosticCode, @@ -247,6 +262,13 @@ function decodeAuthentication(value: unknown): AuthDiagnostics { ...(check.field === undefined ? {} : { field: check.field }), }; }); + if (source.ready) { + if (checks.length !== 1 || checks[0].level !== "info" || checks[0].code !== "auth_ready" + || checks[0].field !== undefined) throw new Error("Workspace API returned invalid authentication diagnostics"); + } else if (!checks.some(({ level }) => level === "error") + || checks.some(({ code }) => code === "auth_ready")) { + throw new Error("Workspace API returned invalid authentication diagnostics"); + } return { ready: source.ready, mode: source.mode as AuthDiagnostics["mode"], checks }; } diff --git a/frontend/src/shell/WorkspaceManager.test.tsx b/frontend/src/shell/WorkspaceManager.test.tsx index 0175650b..c8539c41 100644 --- a/frontend/src/shell/WorkspaceManager.test.tsx +++ b/frontend/src/shell/WorkspaceManager.test.tsx @@ -250,6 +250,29 @@ test("renders one authentication section with configured-group errors and no unm expect(within(section).queryByText(/unmapped/i)).not.toBeInTheDocument(); }); +test("never renders a hostile authentication field rejected by the API decoder", async () => { + const user = userEvent.setup(); + const attacker = "attacker-field-SENTINEL"; + server.use(http.post("/api/workspaces/validate", () => HttpResponse.json({ + workspace, + contract: {}, + activatable: false, + diagnostics: [], + authentication: { + ready: false, + mode: "oidc", + checks: [{ level: "error", code: "oidc_secret_missing", message: "failure", field: attacker }], + }, + }))); + renderManager(); + await user.click(await screen.findByRole("button", { name: "PSD Clinical" })); + + await user.click(screen.getByRole("button", { name: "Validate workspace source" })); + + expect(await screen.findByRole("alert")).toBeVisible(); + expect(screen.queryByText(new RegExp(attacker))).not.toBeInTheDocument(); +}); + test("renders binding_ok as a green connection success", async () => { const user = userEvent.setup(); server.use( diff --git a/frontend/src/shell/WorkspaceManager.tsx b/frontend/src/shell/WorkspaceManager.tsx index 52b807a7..3f547c15 100644 --- a/frontend/src/shell/WorkspaceManager.tsx +++ b/frontend/src/shell/WorkspaceManager.tsx @@ -492,7 +492,7 @@ export function WorkspaceManager({ {authentication.ready ? "Passed" : "Failed"} - {!authentication.ready &&
+ {authentication.checks.some(({ level }) => level === "error") &&
{authentication.checks.filter(({ level }) => level === "error").map(({ code, field, message }) => (

diff --git a/tools/tht/internal/authconfig/commands.go b/tools/tht/internal/authconfig/commands.go index 2bf3bbcc..5143633c 100644 --- a/tools/tht/internal/authconfig/commands.go +++ b/tools/tht/internal/authconfig/commands.go @@ -13,9 +13,13 @@ import ( "net" "net/url" "os" + "os/exec" "path/filepath" "strings" "time" + "unicode" + "unicode/utf16" + "unicode/utf8" "github.com/aritmolab/thothii/tools/tht/internal/compose" "github.com/aritmolab/thothii/tools/tht/internal/config" @@ -76,9 +80,9 @@ func RunWithRunner(ctx context.Context, installation config.Installation, args [ // AuthDiagnostic is the closed JSON contract emitted by the backend diagnostic command. type AuthDiagnostic struct { - Level string `json:"level"` - Code string `json:"code"` - Message string `json:"message"` + Level string `json:"level"` + Code string `json:"code"` + Message string `json:"message"` Field *string `json:"field,omitempty"` } @@ -89,6 +93,19 @@ type AuthDiagnostics struct { Checks []AuthDiagnostic `json:"checks"` } +type authDiagnosticWire struct { + Level string `json:"level"` + Code string `json:"code"` + Message string `json:"message"` + Field json.RawMessage `json:"field"` +} + +type authDiagnosticsWire struct { + Ready bool `json:"ready"` + Mode string `json:"mode"` + Checks []authDiagnosticWire `json:"checks"` +} + func parseCheckArgs(args []string) (jsonMode, interactive bool, err error) { for _, arg := range args { switch arg { @@ -160,20 +177,95 @@ func validAuthDiagnostics(report AuthDiagnostics) bool { if len(report.Checks) == 0 || len(report.Checks) > 129 { return false } + seen := make(map[string]struct{}, len(report.Checks)) + hasError := false for _, check := range report.Checks { - if (check.Level != "error" && check.Level != "info") || check.Message == "" || len(check.Message) > 512 { + if (check.Level != "error" && check.Level != "info") || !safeDiagnosticText(check.Message) { return false } if _, ok := authDiagnosticCodes[check.Code]; !ok { return false } - if check.Field != nil && (*check.Field == "" || len(*check.Field) > 512) { + if check.Field != nil && (!safeDiagnosticText(*check.Field) || + check.Code != "oidc_mapped_group_missing" && check.Code != "oidc_mapped_group_ambiguous") { + return false + } + field := "" + if check.Field != nil { + field = *check.Field + } + key := check.Code + "\x00" + field + if _, duplicate := seen[key]; duplicate { + return false + } + seen[key] = struct{}{} + hasError = hasError || check.Level == "error" + } + if report.Ready { + return len(report.Checks) == 1 && report.Checks[0].Level == "info" && + report.Checks[0].Code == "auth_ready" && report.Checks[0].Field == nil + } + if !hasError { + return false + } + for _, check := range report.Checks { + if check.Code == "auth_ready" { return false } } return true } +func safeDiagnosticText(value string) bool { + if value == "" || utf16Length(value) > 512 || strings.TrimSpace(value) != value { + return false + } + for _, character := range value { + if unicode.IsControl(character) { + return false + } + } + return true +} + +func utf16Length(value string) int { + length := 0 + for _, character := range value { + length += utf16.RuneLen(character) + } + return length +} + +func decodeAuthDiagnostics(value string) (AuthDiagnostics, error) { + if !utf8.ValidString(value) { + return AuthDiagnostics{}, errors.New("authentication diagnostic report is invalid") + } + decoder := json.NewDecoder(strings.NewReader(value)) + decoder.DisallowUnknownFields() + var wire authDiagnosticsWire + if err := decoder.Decode(&wire); err != nil || decoder.Decode(&struct{}{}) != io.EOF { + return AuthDiagnostics{}, errors.New("authentication diagnostic report is invalid") + } + report := AuthDiagnostics{Ready: wire.Ready, Mode: wire.Mode, Checks: make([]AuthDiagnostic, 0, len(wire.Checks))} + for _, item := range wire.Checks { + var field *string + if item.Field != nil { + var decoded string + if string(item.Field) == "null" || json.Unmarshal(item.Field, &decoded) != nil { + return AuthDiagnostics{}, errors.New("authentication diagnostic report is invalid") + } + field = &decoded + } + report.Checks = append(report.Checks, AuthDiagnostic{ + Level: item.Level, Code: item.Code, Message: item.Message, Field: field, + }) + } + if !validAuthDiagnostics(report) { + return AuthDiagnostics{}, errors.New("authentication diagnostic report is invalid") + } + return report, nil +} + func authenticationSecretValues(installation config.Installation) []string { files, err := installation.SecretFiles() if err != nil { @@ -260,18 +352,29 @@ func runCheck(ctx context.Context, installation config.Installation, runner comp if interactive { command = append(command, "--interactive") } - result, err := runner.Run(bounded, installation.ComposeArgs(command...), nil) + result, err := compose.RunBounded(runner, bounded, installation.ComposeArgs(command...), nil, compose.CaptureLimits{ + StdoutBytes: maxAuthDiagnosticOutputBytes, + StderrBytes: maxAuthDiagnosticOutputBytes, + }) secrets := authenticationSecretValues(installation) - if err != nil || result.ExitCode != 0 || len(result.Stdout) > maxAuthDiagnosticOutputBytes || len(result.Stderr) > maxAuthDiagnosticOutputBytes { + var exitError *exec.ExitError + validProcessOutcome := result.ExitCode == 0 && err == nil || + result.ExitCode == 1 && (err == nil || errors.As(err, &exitError)) + if !validProcessOutcome { return AuthDiagnostics{}, "", errors.New("authentication diagnostic command failed") } - decoder := json.NewDecoder(strings.NewReader(result.Stdout)) - decoder.DisallowUnknownFields() - var report AuthDiagnostics - if err := decoder.Decode(&report); err != nil || decoder.Decode(&struct{}{}) != io.EOF || !validAuthDiagnostics(report) { + report, decodeErr := decodeAuthDiagnostics(result.Stdout) + if decodeErr != nil { return AuthDiagnostics{}, "", errors.New("authentication diagnostic report is invalid") } - return sanitizeAuthDiagnostics(report, secrets), devicePrompt(result.Stderr, secrets), nil + if report.Ready != (result.ExitCode == 0) { + return AuthDiagnostics{}, "", errors.New("authentication diagnostic command failed") + } + safe := sanitizeAuthDiagnostics(report, secrets) + if !validAuthDiagnostics(safe) { + return AuthDiagnostics{}, "", errors.New("authentication diagnostic report is invalid") + } + return safe, devicePrompt(result.Stderr, secrets), nil } func authFailure(stderr io.Writer, message string) int { diff --git a/tools/tht/internal/authconfig/commands_test.go b/tools/tht/internal/authconfig/commands_test.go index 3c194aeb..2c83dc4e 100644 --- a/tools/tht/internal/authconfig/commands_test.go +++ b/tools/tht/internal/authconfig/commands_test.go @@ -49,6 +49,113 @@ func TestAuthCheckRunsOneShotCoreDiagnosticWithPristineJSON(t *testing.T) { } } +func TestAuthCheckEmitsValidFailedReportFromRealExitError(t *testing.T) { + installation := authInstallation(newAuthDirectory(t)) + runner := compose.NewRunner(writeAuthExecutable(t, `#!/bin/sh +printf '%s\n' '{"ready":false,"mode":"oidc","checks":[{"level":"error","code":"oidc_secret_missing","message":"A required OIDC or group catalog secret is unavailable."}]}' +exit 1 +`)) + var stdout, stderr bytes.Buffer + + code := RunWithRunner(context.Background(), installation, []string{"check", "--json"}, strings.NewReader(""), &stdout, &stderr, runner) + + if code != 1 { + t.Fatalf("auth check = %d, want diagnostic failure 1; stdout=%q stderr=%q", code, stdout.String(), stderr.String()) + } + var report AuthDiagnostics + if err := json.Unmarshal(stdout.Bytes(), &report); err != nil || report.Ready || report.Checks[0].Code != "oidc_secret_missing" { + t.Fatalf("auth check stdout is not the pristine failed report: %q: %#v, %v", stdout.String(), report, err) + } + if stderr.Len() != 0 { + t.Fatalf("auth check stderr = %q, want empty", stderr.String()) + } +} + +func TestAuthCheckRejectsExitAndReportSemanticMismatches(t *testing.T) { + ready := `{"ready":true,"mode":"oidc","checks":[{"level":"info","code":"auth_ready","message":"Authentication is ready."}]}` + failed := `{"ready":false,"mode":"oidc","checks":[{"level":"error","code":"oidc_secret_missing","message":"A required OIDC or group catalog secret is unavailable."}]}` + for _, test := range []struct { + name string + exit int + report string + }{ + {name: "zero with failed report", exit: 0, report: failed}, + {name: "one with ready report", exit: 1, report: ready}, + {name: "two with failed report", exit: 2, report: failed}, + } { + t.Run(test.name, func(t *testing.T) { + runner := runnerFunc(func(_ context.Context, _ []string, _ io.Reader) (compose.Result, error) { + var err error + if test.exit != 0 { + err = errors.New("process exited") + } + return compose.Result{Stdout: test.report, ExitCode: test.exit}, err + }) + var stdout, stderr bytes.Buffer + code := RunWithRunner(context.Background(), authInstallation(newAuthDirectory(t)), []string{"check", "--json"}, strings.NewReader(""), &stdout, &stderr, runner) + if code != 1 || stdout.Len() != 0 || stderr.String() != "tht: authentication diagnostics could not be completed\n" { + t.Fatalf("mismatch accepted: code=%d stdout=%q stderr=%q", code, stdout.String(), stderr.String()) + } + }) + } +} + +func TestAuthCheckRejectsCancellationEvenIfTheKilledChildReportsExitOne(t *testing.T) { + failed := `{"ready":false,"mode":"oidc","checks":[{"level":"error","code":"oidc_secret_missing","message":"A required OIDC or group catalog secret is unavailable."}]}` + runner := runnerFunc(func(_ context.Context, _ []string, _ io.Reader) (compose.Result, error) { + return compose.Result{Stdout: failed, ExitCode: 1}, context.DeadlineExceeded + }) + var stdout, stderr bytes.Buffer + + code := RunWithRunner(context.Background(), authInstallation(newAuthDirectory(t)), []string{"check", "--json"}, strings.NewReader(""), &stdout, &stderr, runner) + + if code != 1 || stdout.Len() != 0 || stderr.String() != "tht: authentication diagnostics could not be completed\n" { + t.Fatalf("cancelled report accepted: code=%d stdout=%q stderr=%q", code, stdout.String(), stderr.String()) + } +} + +func TestAuthDiagnosticsContractRejectsContradictionsDuplicatesAndAttackerFields(t *testing.T) { + field := "Configured Group" + validFailure := AuthDiagnostics{Ready: false, Mode: "oidc", Checks: []AuthDiagnostic{{ + Level: "error", Code: "oidc_mapped_group_missing", Message: "A configured authorization group does not exist.", Field: &field, + }}} + if !validAuthDiagnostics(validFailure) { + t.Fatal("valid failed report was rejected") + } + for _, report := range []AuthDiagnostics{ + {Ready: true, Mode: "oidc", Checks: []AuthDiagnostic{{Level: "error", Code: "oidc_secret_missing", Message: "failure"}}}, + {Ready: false, Mode: "oidc", Checks: []AuthDiagnostic{{Level: "info", Code: "auth_ready", Message: "ready"}}}, + {Ready: false, Mode: "oidc", Checks: []AuthDiagnostic{{Level: "info", Code: "auth_config_invalid", Message: "not an error"}}}, + {Ready: false, Mode: "oidc", Checks: []AuthDiagnostic{ + {Level: "error", Code: "oidc_secret_missing", Message: "failure"}, + {Level: "error", Code: "oidc_secret_missing", Message: "duplicate"}, + }}, + {Ready: false, Mode: "oidc", Checks: []AuthDiagnostic{{Level: "error", Code: "oidc_secret_missing", Message: "failure", Field: &field}}}, + {Ready: false, Mode: "oidc", Checks: []AuthDiagnostic{{Level: "error", Code: "oidc_mapped_group_missing", Message: "failure", Field: stringPointer(" attacker ")}}}, + {Ready: false, Mode: "oidc", Checks: []AuthDiagnostic{{Level: "error", Code: "oidc_mapped_group_missing", Message: "failure", Field: stringPointer("attacker\u0085field")}}}, + } { + if validAuthDiagnostics(report) { + t.Fatalf("invalid authentication report accepted: %#v", report) + } + } +} + +func TestAuthCheckRejectsNullAndUnexpectedDiagnosticFields(t *testing.T) { + for _, report := range []string{ + `{"ready":false,"mode":"oidc","checks":[{"level":"error","code":"oidc_mapped_group_missing","message":"failure","field":null}]}`, + `{"ready":false,"mode":"oidc","checks":[{"level":"error","code":"oidc_secret_missing","message":"failure","unexpected":"attacker"}]}`, + } { + runner := runnerFunc(func(_ context.Context, _ []string, _ io.Reader) (compose.Result, error) { + return compose.Result{Stdout: report, ExitCode: 1}, nil + }) + var stdout, stderr bytes.Buffer + code := RunWithRunner(context.Background(), authInstallation(newAuthDirectory(t)), []string{"check", "--json"}, strings.NewReader(""), &stdout, &stderr, runner) + if code != 1 || stdout.Len() != 0 || stderr.String() != "tht: authentication diagnostics could not be completed\n" { + t.Fatalf("hostile field accepted: code=%d stdout=%q stderr=%q", code, stdout.String(), stderr.String()) + } + } +} + func TestAuthCheckInteractiveForwardsOnlyTheValidatedDevicePrompt(t *testing.T) { installation := authInstallation(newAuthDirectory(t)) var calls [][]string @@ -96,8 +203,8 @@ func TestAuthCheckRedactsFailedCoreOutputAndRejectsMalformedReports(t *testing.T } runner := runnerFunc(func(_ context.Context, _ []string, _ io.Reader) (compose.Result, error) { return compose.Result{ - Stdout: "not-json auth-check-secret token-sentinel /private/sentinel $argon2id$hash-sentinel", - Stderr: "auth-check-secret cookie-sentinel /private/sentinel", + Stdout: "not-json auth-check-secret token-sentinel /private/sentinel $argon2id$hash-sentinel", + Stderr: "auth-check-secret cookie-sentinel /private/sentinel", ExitCode: 23, }, errors.New("core failed") }) @@ -419,6 +526,17 @@ func writePasswordFile(t *testing.T, password string) string { return path } +func writeAuthExecutable(t *testing.T, contents string) string { + t.Helper() + path := filepath.Join(t.TempDir(), "fake-docker") + if err := os.WriteFile(path, []byte(contents), 0o700); err != nil { + t.Fatal(err) + } + return path +} + +func stringPointer(value string) *string { return &value } + func authInstallation(directory string) config.Installation { installation := config.Installation{} installation.Authentication.ConfigDirectory = directory diff --git a/tools/tht/internal/compose/process_unix.go b/tools/tht/internal/compose/process_unix.go new file mode 100644 index 00000000..7688d1a1 --- /dev/null +++ b/tools/tht/internal/compose/process_unix.go @@ -0,0 +1,37 @@ +//go:build !windows + +package compose + +import ( + "errors" + "os/exec" + "syscall" + "time" +) + +const gracefulTerminationBound = 500 * time.Millisecond +const finalTerminationBound = 2 * time.Second + +func configureProcess(command *exec.Cmd) { + command.SysProcAttr = &syscall.SysProcAttr{Setpgid: true} +} + +func terminateProcess(command *exec.Cmd, done <-chan error) error { + if command.Process == nil { + return nil + } + _ = syscall.Kill(-command.Process.Pid, syscall.SIGINT) + select { + case <-done: + return nil + case <-time.After(gracefulTerminationBound): + } + _ = syscall.Kill(-command.Process.Pid, syscall.SIGKILL) + _ = command.Process.Kill() + select { + case <-done: + return nil + case <-time.After(finalTerminationBound): + return errors.New("Docker command could not be reaped") + } +} diff --git a/tools/tht/internal/compose/process_windows.go b/tools/tht/internal/compose/process_windows.go new file mode 100644 index 00000000..1e083ee2 --- /dev/null +++ b/tools/tht/internal/compose/process_windows.go @@ -0,0 +1,30 @@ +//go:build windows + +package compose + +import ( + "errors" + "os/exec" + "syscall" + "time" +) + +const finalTerminationBound = 2 * time.Second +const createNewProcessGroup = 0x00000200 + +func configureProcess(command *exec.Cmd) { + command.SysProcAttr = &syscall.SysProcAttr{CreationFlags: createNewProcessGroup} +} + +func terminateProcess(command *exec.Cmd, done <-chan error) error { + if command.Process == nil { + return nil + } + _ = command.Process.Kill() + select { + case <-done: + return nil + case <-time.After(finalTerminationBound): + return errors.New("Docker command could not be reaped") + } +} diff --git a/tools/tht/internal/compose/runner.go b/tools/tht/internal/compose/runner.go index aeb85860..17324622 100644 --- a/tools/tht/internal/compose/runner.go +++ b/tools/tht/internal/compose/runner.go @@ -2,17 +2,29 @@ package compose import ( - "bytes" "context" "errors" "fmt" "io" "os" "os/exec" + "sync" "github.com/aritmolab/thothii/tools/tht/internal/config" ) +const defaultCaptureBytes = 4 * 1024 * 1024 +const maximumCaptureBytes = 64 * 1024 * 1024 + +// ErrOutputLimit reports that a child exceeded one of its capture limits. +var ErrOutputLimit = errors.New("Docker command output limit exceeded") + +// CaptureLimits bounds each captured stream while the child is running. +type CaptureLimits struct { + StdoutBytes int + StderrBytes int +} + // Result is the captured output and process exit code for one Docker invocation. type Result struct { Stdout string @@ -25,6 +37,10 @@ type Runner interface { Run(context.Context, []string, io.Reader) (Result, error) } +type boundedRunner interface { + RunBounded(context.Context, []string, io.Reader, CaptureLimits) (Result, error) +} + // execRunner executes the Docker CLI. It never invokes a shell. type execRunner struct { binary string @@ -40,21 +56,101 @@ func NewRunner(binary string) Runner { // Run invokes Docker with the supplied argument array and optional standard input. func (r execRunner) Run(ctx context.Context, args []string, stdin io.Reader) (Result, error) { - command := exec.CommandContext(ctx, r.binary, args...) + return r.RunBounded(ctx, args, stdin, CaptureLimits{ + StdoutBytes: defaultCaptureBytes, + StderrBytes: defaultCaptureBytes, + }) +} + +// RunBounded invokes Docker while enforcing both stream limits during capture. +func (r execRunner) RunBounded(ctx context.Context, args []string, stdin io.Reader, limits CaptureLimits) (Result, error) { + if err := validCaptureLimits(limits); err != nil { + return Result{}, err + } + if err := ctx.Err(); err != nil { + return Result{}, err + } + command := exec.Command(r.binary, args...) + configureProcess(command) command.Stdin = stdin - var stdout, stderr bytes.Buffer - command.Stdout = &stdout - command.Stderr = &stderr - err := command.Run() + overflow := make(chan struct{}, 1) + stdout := newCappedBuffer(limits.StdoutBytes, overflow) + stderr := newCappedBuffer(limits.StderrBytes, overflow) + command.Stdout = stdout + command.Stderr = stderr + if err := command.Start(); err != nil { + return startFailure(err) + } + done := make(chan error, 1) + go func() { done <- command.Wait() }() + var err error + select { + case err = <-done: + case <-ctx.Done(): + _ = terminateProcess(command, done) + err = ctx.Err() + case <-overflow: + _ = terminateProcess(command, done) + err = ErrOutputLimit + } result := Result{Stdout: stdout.String(), Stderr: stderr.String()} + if command.ProcessState != nil { + result.ExitCode = command.ProcessState.ExitCode() + } + if stdout.Overflowed() || stderr.Overflowed() { + return result, ErrOutputLimit + } if err == nil { return result, nil } + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) || errors.Is(err, ErrOutputLimit) { + return result, err + } var exitError *exec.ExitError if errors.As(err, &exitError) { result.ExitCode = exitError.ExitCode() return result, err } + return result, err +} + +// RunBounded uses the production runner's during-capture limits while retaining compatibility +// with injected runners, whose already-bounded test results are checked before use. +func RunBounded(runner Runner, ctx context.Context, args []string, stdin io.Reader, limits CaptureLimits) (Result, error) { + if err := validCaptureLimits(limits); err != nil { + return Result{}, err + } + if bounded, ok := runner.(boundedRunner); ok { + return bounded.RunBounded(ctx, args, stdin, limits) + } + result, err := runner.Run(ctx, args, stdin) + if len(result.Stdout) > limits.StdoutBytes || len(result.Stderr) > limits.StderrBytes { + return Result{ + Stdout: boundedString(result.Stdout, limits.StdoutBytes), + Stderr: boundedString(result.Stderr, limits.StderrBytes), + ExitCode: result.ExitCode, + }, ErrOutputLimit + } + return result, err +} + +func validCaptureLimits(limits CaptureLimits) error { + if limits.StdoutBytes < 1 || limits.StderrBytes < 1 || + limits.StdoutBytes > maximumCaptureBytes || limits.StderrBytes > maximumCaptureBytes { + return errors.New("Docker command capture limits are invalid") + } + return nil +} + +func boundedString(value string, maximum int) string { + if len(value) <= maximum { + return value + } + return value[:maximum] +} + +func startFailure(err error) (Result, error) { + result := Result{} if errors.Is(err, exec.ErrNotFound) || errors.Is(err, os.ErrNotExist) { result.ExitCode = 127 return result, fmt.Errorf("%w: %w", exec.ErrNotFound, err) @@ -62,6 +158,55 @@ func (r execRunner) Run(ctx context.Context, args []string, stdin io.Reader) (Re return result, err } +type cappedBuffer struct { + mu sync.Mutex + contents []byte + maximum int + overflow chan<- struct{} + exceeded bool +} + +func newCappedBuffer(maximum int, overflow chan<- struct{}) *cappedBuffer { + capacity := maximum + if capacity > 4096 { + capacity = 4096 + } + return &cappedBuffer{contents: make([]byte, 0, capacity), maximum: maximum, overflow: overflow} +} + +func (b *cappedBuffer) Write(value []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + remaining := b.maximum - len(b.contents) + if remaining > 0 { + kept := len(value) + if kept > remaining { + kept = remaining + } + b.contents = append(b.contents, value[:kept]...) + } + if len(value) > remaining && !b.exceeded { + b.exceeded = true + select { + case b.overflow <- struct{}{}: + default: + } + } + return len(value), nil +} + +func (b *cappedBuffer) String() string { + b.mu.Lock() + defer b.mu.Unlock() + return string(b.contents) +} + +func (b *cappedBuffer) Overflowed() bool { + b.mu.Lock() + defer b.mu.Unlock() + return b.exceeded +} + // InstallationRunner applies an installation's validated Compose arguments to commands that // explicitly start with compose. Direct Docker image commands remain host-side. type InstallationRunner struct { @@ -84,3 +229,11 @@ func (r InstallationRunner) Run(ctx context.Context, args []string, stdin io.Rea } return r.Runner.Run(ctx, args, stdin) } + +// RunBounded preserves bounded capture when Compose argument injection is wrapped per installation. +func (r InstallationRunner) RunBounded(ctx context.Context, args []string, stdin io.Reader, limits CaptureLimits) (Result, error) { + if len(args) > 0 && args[0] == "compose" { + return RunBounded(r.Runner, ctx, r.Installation.ComposeArgs(args[1:]...), stdin, limits) + } + return RunBounded(r.Runner, ctx, args, stdin, limits) +} diff --git a/tools/tht/internal/compose/runner_test.go b/tools/tht/internal/compose/runner_test.go index 45697240..7328779c 100644 --- a/tools/tht/internal/compose/runner_test.go +++ b/tools/tht/internal/compose/runner_test.go @@ -9,6 +9,7 @@ import ( "path/filepath" "strings" "testing" + "time" "github.com/aritmolab/thothii/tools/tht/internal/config" ) @@ -30,6 +31,45 @@ func TestRunnerPassesEachArgumentWithoutShellSplitting(t *testing.T) { } } +func TestRunnerBoundsFloodingOutputDuringCapture(t *testing.T) { + t.Parallel() + + runner := NewRunner(writeExecutable(t, "#!/bin/sh\nwhile :; do printf '0123456789abcdef'; printf 'fedcba9876543210' >&2; done\n")) + started := time.Now() + result, err := RunBounded(runner, context.Background(), []string{"compose", "run", "--rm", "core"}, nil, CaptureLimits{ + StdoutBytes: 1024, + StderrBytes: 1024, + }) + if !errors.Is(err, ErrOutputLimit) { + t.Fatalf("RunBounded() error = %v, want ErrOutputLimit", err) + } + if len(result.Stdout) > 1024 || len(result.Stderr) > 1024 { + t.Fatalf("captured output exceeded limits: stdout=%d stderr=%d", len(result.Stdout), len(result.Stderr)) + } + if elapsed := time.Since(started); elapsed > 5*time.Second { + t.Fatalf("overflow teardown took %s", elapsed) + } +} + +func TestRunnerCancelsAndReapsAHangingChildWithinFinalBound(t *testing.T) { + t.Parallel() + + runner := NewRunner(writeExecutable(t, "#!/bin/sh\ntrap '' TERM INT\nwhile :; do sleep 1; done\n")) + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + started := time.Now() + _, err := RunBounded(runner, ctx, []string{"compose", "run", "--rm", "core"}, nil, CaptureLimits{ + StdoutBytes: 1024, + StderrBytes: 1024, + }) + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("RunBounded() error = %v, want deadline exceeded", err) + } + if elapsed := time.Since(started); elapsed > 5*time.Second { + t.Fatalf("cancellation teardown took %s", elapsed) + } +} + func TestRunnerReturnsTheChildExitCode(t *testing.T) { t.Parallel() diff --git a/tools/tht/internal/doctor/report.go b/tools/tht/internal/doctor/report.go index c6199e5e..c23ccef7 100644 --- a/tools/tht/internal/doctor/report.go +++ b/tools/tht/internal/doctor/report.go @@ -155,27 +155,36 @@ func RunWithProbe(ctx context.Context, installation config.Installation, runner status, statusAvailable, servicesCheck := serviceStatus(ctx, installation, runner, secretValues) coreRunning := false + coreHealthy := false if statusAvailable { - var err error - coreRunning, err = service.CoreRunning(status) - if err != nil { + var runningErr, healthyErr error + coreRunning, runningErr = service.CoreRunning(status) + coreHealthy, healthyErr = service.CoreHealthy(status) + if runningErr != nil || healthyErr != nil { coreRunning = false + coreHealthy = false } } if !configReady || !coreRunning { add("authentication", StatusSkipped, "core is unavailable") + } else if !coreHealthy { + add("authentication", StatusFailed, "core is running but unhealthy") } else if authenticationCheck(ctx, installation, runner, secretValues) { add("authentication", StatusPassed, "container-local authentication diagnostics passed") } else { add("authentication", StatusFailed, "container-local authentication diagnostics failed") } add(servicesCheck.Name, servicesCheck.Status, servicesCheck.Detail) - if !coreRunning { - add("core-http", StatusSkipped, "core is not running") - add("frontend-http", StatusSkipped, "core is not running") - add("workspace-registry", StatusSkipped, "core is not running") - add("workflow", StatusSkipped, "core is not running") - add("pi", StatusSkipped, "core is not running") + if !coreHealthy || servicesCheck.Status != StatusPassed { + detail := "required services are not healthy" + if !coreRunning { + detail = "core is not running" + } + add("core-http", StatusSkipped, detail) + add("frontend-http", StatusSkipped, detail) + add("workspace-registry", StatusSkipped, detail) + add("workflow", StatusSkipped, detail) + add("pi", StatusSkipped, detail) return finalize(report), nil } diff --git a/tools/tht/internal/doctor/report_test.go b/tools/tht/internal/doctor/report_test.go index cee51175..15e1f73d 100644 --- a/tools/tht/internal/doctor/report_test.go +++ b/tools/tht/internal/doctor/report_test.go @@ -60,6 +60,43 @@ func TestRunSkipsContainerDiagnosticsWhenCoreIsStopped(t *testing.T) { } } +func TestRunFailsAuthenticationWithoutExecWhenCoreIsRunningButUnhealthy(t *testing.T) { + installation := doctorInstallation(t, "") + runner := &doctorRunner{services: unhealthyCoreServices} + + report, err := Run(context.Background(), installation, runner) + if err != nil { + t.Fatal(err) + } + if report.OK || checkStatus(report, "authentication") != StatusFailed { + t.Fatalf("Run() report = %#v, want deterministic failed authentication", report) + } + assertChecklist(t, report, []string{"descriptor", "files", "docker", "compose", "configuration", "authentication", "services", "core-http", "frontend-http", "workspace-registry", "workflow", "pi"}) + if strings.Contains(strings.Join(runner.calls, "\n"), " exec -T ") { + t.Fatalf("Run() invoked exec -T while core was unhealthy: %v", runner.calls) + } + if detail := checkDetail(report, "authentication"); detail != "core is running but unhealthy" { + t.Fatalf("authentication detail = %q, want deterministic unhealthy detail", detail) + } +} + +func TestRunExecutesOnlyAuthenticationWhenCoreIsHealthyButAnotherServiceIsUnhealthy(t *testing.T) { + installation := doctorInstallation(t, "") + runner := &doctorRunner{services: unhealthyFrontendServices} + + report, err := Run(context.Background(), installation, runner) + if err != nil { + t.Fatal(err) + } + if checkStatus(report, "authentication") != StatusPassed || checkStatus(report, "services") != StatusFailed { + t.Fatalf("Run() report = %#v, want auth passed before failed services", report) + } + calls := strings.Join(runner.calls, "\n") + if strings.Count(calls, " exec -T ") != 1 || !strings.Contains(calls, "exec -T core node dist/auth/diagnostic-command.js --json") { + t.Fatalf("Run() calls = %s, want only the healthy-core authentication exec", calls) + } +} + // Catches host-Python diagnostics or omission of workflow/Pi checks once core is healthy. func TestRunUsesOnlyContainerLocalWorkflowAndPiDiagnosticsWhenCoreRuns(t *testing.T) { installation := doctorInstallation(t, "") @@ -240,6 +277,15 @@ func checkStatus(report Report, name string) string { return "" } +func checkDetail(report Report, name string) string { + for _, check := range report.Checks { + if check.Name == name { + return check.Detail + } + } + return "" +} + func reportText(report Report) string { parts := make([]string, 0, len(report.Checks)) for _, check := range report.Checks { @@ -273,3 +319,19 @@ const stoppedServices = `[ {"Service":"core","State":"exited","Health":""}, {"Service":"frontend","State":"running","Health":"healthy"} ]` + +const unhealthyCoreServices = `[ + {"Service":"core","State":"running","Health":"unhealthy"}, + {"Service":"frontend","State":"running","Health":"healthy"}, + {"Service":"qdrant","State":"running","Health":"healthy"}, + {"Service":"embedding","State":"running","Health":"healthy"}, + {"Service":"embedding-model-init","State":"exited","ExitCode":0} +]` + +const unhealthyFrontendServices = `[ + {"Service":"core","State":"running","Health":"healthy"}, + {"Service":"frontend","State":"running","Health":"unhealthy"}, + {"Service":"qdrant","State":"running","Health":"healthy"}, + {"Service":"embedding","State":"running","Health":"healthy"}, + {"Service":"embedding-model-init","State":"exited","ExitCode":0} +]` diff --git a/tools/tht/internal/service/service.go b/tools/tht/internal/service/service.go index cf8494a2..a1d9d8c9 100644 --- a/tools/tht/internal/service/service.go +++ b/tools/tht/internal/service/service.go @@ -97,6 +97,20 @@ func CoreRunning(value string) (bool, error) { return false, nil } +// CoreHealthy reports whether core is both running and certified healthy by Compose. +func CoreHealthy(value string) (bool, error) { + statuses, err := parseStatuses(value) + if err != nil { + return false, err + } + for _, item := range statuses { + if item.Service == "core" { + return strings.EqualFold(item.State, "running") && strings.EqualFold(item.Health, "healthy"), nil + } + } + return false, nil +} + // Healthy verifies the full expected Compose service set, including the one-shot model initializer. func Healthy(value string) error { statuses, err := parseStatuses(value)