From 6b828288e3e2860329c8afd7fa076d3f6f2ff894 Mon Sep 17 00:00:00 2001 From: mptyl Date: Wed, 5 Aug 2026 01:35:42 +0200 Subject: [PATCH] fix: preserve custom provider smoke config --- backend/src/pi/provider-smoke.ts | 100 ++++++++++++++++++++++++- backend/test/pi-provider-smoke.test.ts | 80 ++++++++++++++++++++ 2 files changed, 176 insertions(+), 4 deletions(-) diff --git a/backend/src/pi/provider-smoke.ts b/backend/src/pi/provider-smoke.ts index fc8400d1..57c3c8d8 100644 --- a/backend/src/pi/provider-smoke.ts +++ b/backend/src/pi/provider-smoke.ts @@ -1,5 +1,8 @@ import { spawn as nodeSpawn, type ChildProcessWithoutNullStreams } from "node:child_process"; -import { mkdirSync, mkdtempSync, readFileSync, rmSync, writeFileSync } from "node:fs"; +import { + closeSync, constants, fstatSync, lstatSync, mkdirSync, mkdtempSync, openSync, + readFileSync, rmSync, writeFileSync, +} from "node:fs"; import { homedir, tmpdir } from "node:os"; import { join } from "node:path"; import type { AppConfig } from "../config.js"; @@ -11,6 +14,7 @@ import { buildPiChildEnv, canonicalPiProvider } from "./provider-credentials.js" import type { PiReasoning } from "./management.js"; const SMOKE_PROMPT = "Provider health check. Reply with exactly OK."; +const MAX_AGENT_CONFIG_BYTES = 1024 * 1024; const SMOKE_ARGS = [ "--mode", "rpc", "--no-session", @@ -40,6 +44,7 @@ interface ProviderSmokeOptions { ) => ChildProcessWithoutNullStreams; authProviders?: () => ReadonlySet; readAuthStore?: () => string; + readModelsStore?: () => string | undefined; } export function createPiProviderSmoke( @@ -77,11 +82,26 @@ export function createPiProviderSmoke( mkdirSync(isolatedAgentDir, { mode: 0o700 }); if (configuredAuthProviders.has(canonicalProvider)) { const authStore = selectedProviderAuthStore( - options.readAuthStore?.() ?? readConfiguredAuthStore(), + options.readAuthStore?.() ?? readConfiguredAgentFile("auth.json"), canonicalProvider, ); writeFileSync(join(isolatedAgentDir, "auth.json"), authStore, { mode: 0o600, flag: "wx" }); } + const configuredModels = options.readModelsStore + ? options.readModelsStore() + : readConfiguredAgentFile("models.json", true); + if (configuredModels !== undefined) { + const modelsStore = selectedProviderModelsStore( + configuredModels, + canonicalProvider, + model, + ); + if (modelsStore !== undefined) { + writeFileSync(join(isolatedAgentDir, "models.json"), modelsStore, { + mode: 0o600, flag: "wx", + }); + } + } env.PI_CODING_AGENT_DIR = isolatedAgentDir; child = spawnFn(config.piBin, [...SMOKE_ARGS], { cwd: isolatedCwd, env }); @@ -147,9 +167,35 @@ function messageUsesTool(message: any): boolean { && message.content.some((content: any) => content?.type === "toolCall")); } -function readConfiguredAuthStore(): string { +function readConfiguredAgentFile(name: "auth.json"): string; +function readConfiguredAgentFile(name: "models.json", optional: true): string | undefined; +function readConfiguredAgentFile( + name: "auth.json" | "models.json", + optional = false, +): string | undefined { const configuredAgentDir = process.env.PI_CODING_AGENT_DIR ?? join(homedir(), ".pi", "agent"); - return readFileSync(join(configuredAgentDir, "auth.json"), "utf8"); + const path = join(configuredAgentDir, name); + let fd: number | undefined; + try { + const before = lstatSync(path); + if (!before.isFile() || before.isSymbolicLink() || before.size > MAX_AGENT_CONFIG_BYTES) { + throw providerFailure(); + } + fd = openSync(path, constants.O_RDONLY | constants.O_NOFOLLOW); + const opened = fstatSync(fd); + if (!opened.isFile() || opened.size > MAX_AGENT_CONFIG_BYTES + || before.dev !== opened.dev || before.ino !== opened.ino) { + throw providerFailure(); + } + return readFileSync(fd, "utf8"); + } catch (error) { + if (optional && (error as NodeJS.ErrnoException)?.code === "ENOENT") return undefined; + throw providerFailure(); + } finally { + if (fd !== undefined) { + try { closeSync(fd); } catch { /* preserve the sanitized smoke outcome */ } + } + } } function selectedProviderAuthStore(raw: string, provider: string): string { @@ -161,6 +207,52 @@ function selectedProviderAuthStore(raw: string, provider: string): string { return JSON.stringify({ [entry[0]]: entry[1] }); } +const PROVIDER_CONFIG_FIELDS = [ + "name", "baseUrl", "apiKey", "api", "headers", "compat", "authHeader", +] as const; + +function selectedProviderModelsStore(raw: string, provider: string, model: string): string | undefined { + const parsed: unknown = JSON.parse(raw); + if (!parsed || typeof parsed !== "object" || Array.isArray(parsed)) throw providerFailure(); + const providers = (parsed as { providers?: unknown }).providers; + if (!providers || typeof providers !== "object" || Array.isArray(providers)) { + throw providerFailure(); + } + const entry = Object.entries(providers as Record) + .find(([key]) => key.trim().toLowerCase() === provider); + if (!entry) return undefined; + const providerConfig = entry[1]; + if (!providerConfig || typeof providerConfig !== "object" || Array.isArray(providerConfig)) { + throw providerFailure(); + } + const source = providerConfig as Record; + const selected: Record = {}; + for (const field of PROVIDER_CONFIG_FIELDS) { + if (Object.hasOwn(source, field)) selected[field] = source[field]; + } + if (Object.hasOwn(source, "models")) { + if (!Array.isArray(source.models)) throw providerFailure(); + let selectedModel: unknown; + for (const candidate of source.models) { + if (candidate && typeof candidate === "object" && !Array.isArray(candidate) + && (candidate as { id?: unknown }).id === model) { + selectedModel = candidate; + } + } + if (selectedModel !== undefined) selected.models = [selectedModel]; + } + if (Object.hasOwn(source, "modelOverrides")) { + const overrides = source.modelOverrides; + if (!overrides || typeof overrides !== "object" || Array.isArray(overrides)) { + throw providerFailure(); + } + if (Object.hasOwn(overrides, model)) { + selected.modelOverrides = { [model]: (overrides as Record)[model] }; + } + } + return JSON.stringify({ providers: { [entry[0]]: selected } }); +} + function waitForProviderTurn(rpc: RpcClient, child: ChildProcessWithoutNullStreams): Promise { return new Promise((resolve, reject) => { let failed = false; diff --git a/backend/test/pi-provider-smoke.test.ts b/backend/test/pi-provider-smoke.test.ts index d2ac069a..08260a8c 100644 --- a/backend/test/pi-provider-smoke.test.ts +++ b/backend/test/pi-provider-smoke.test.ts @@ -24,6 +24,83 @@ function rpcChild(onCommand: (command: any, emit: (message: unknown) => void) => return child; } +// Catches an isolated smoke agent that copies auth.json but drops the selected custom +// provider/model from models.json, causing set_model to fail before the real request. +test("provider smoke reaches the selected custom provider from an isolated models.json", async () => { + let isolatedAgentDir: string | undefined; + let providerRequests = 0; + const child = rpcChild((command, emit) => { + if (command.type === "set_model") { + const models = JSON.parse(readFileSync(`${isolatedAgentDir}/models.json`, "utf8")); + const selectedProvider = models.providers?.[command.provider]; + const selectedModel = selectedProvider?.models?.find( + (candidate: { id?: unknown }) => candidate.id === command.modelId, + ); + emit({ type: "response", id: command.id, success: Boolean(selectedModel) }); + } + if (command.type === "set_thinking_level") { + emit({ type: "response", id: command.id, success: true }); + } + if (command.type === "prompt") { + providerRequests++; + emit({ + type: "message_end", + message: { role: "assistant", stopReason: "stop", content: "must-not-be-returned" }, + }); + emit({ + type: "agent_end", + messages: [{ role: "assistant", stopReason: "stop", content: "must-not-be-returned" }], + }); + } + }); + const smoke = createPiProviderSmoke(loadConfig({ PI_BIN: "/usr/local/bin/pi" }), { + spawnFn: (_command, _args, options) => { + isolatedAgentDir = options.env.PI_CODING_AGENT_DIR; + expect(readdirSync(isolatedAgentDir)).toEqual(["auth.json", "models.json"]); + expect(JSON.parse(readFileSync(`${isolatedAgentDir}/models.json`, "utf8"))).toEqual({ + providers: { + "custom-openai": { + baseUrl: "https://selected.invalid/v1", + api: "openai-completions", + models: [{ id: "selected-model", name: "Selected model", reasoning: true }], + }, + }, + }); + return child; + }, + authProviders: () => new Set(["custom-openai"]), + readAuthStore: () => JSON.stringify({ + "custom-openai": { type: "api_key", key: "test-only" }, + unrelated: { type: "api_key", key: "must-not-enter-isolated-context" }, + }), + readModelsStore: () => JSON.stringify({ + providers: { + "custom-openai": { + baseUrl: "https://selected.invalid/v1", + api: "openai-completions", + models: [ + { id: "selected-model", name: "Selected model", reasoning: true }, + { id: "unrelated-model", name: "Must not enter isolated context" }, + ], + }, + unrelated: { + baseUrl: "https://unrelated.invalid/v1", + api: "openai-completions", + apiKey: "!must-not-run", + models: [{ id: "unrelated-model" }], + }, + }, + }), + }); + + await expect(smoke({ + provider: "custom-openai", model: "selected-model", reasoning: "medium", timeoutMs: 750, + })).resolves.toBeUndefined(); + expect(providerRequests).toBe(1); + expect(child.kill).toHaveBeenCalledOnce(); + expect(isolatedAgentDir && existsSync(dirname(isolatedAgentDir))).toBe(false); +}); + // Catches a provider smoke process that runs from the trusted harness or leaves Pi tools, // extensions, skills, context files, templates, themes, or session persistence enabled. test("provider smoke makes one configured request from an isolated no-capability Pi process", async () => { @@ -71,6 +148,7 @@ test("provider smoke makes one configured request from an isolated no-capability zai: { type: "api_key", key: "test-only" }, deepseek: { type: "api_key", key: "must-not-enter-isolated-context" }, }), + readModelsStore: () => undefined, }); await expect(smoke({ @@ -144,6 +222,7 @@ test.each(unexpectedToolEvents)("provider smoke fails closed on an unexpected $n spawnFn: () => child, authProviders: () => new Set(["zai"]), readAuthStore: () => '{"zai":{"type":"api_key","key":"test-only"}}', + readModelsStore: () => undefined, }); await expect(smoke({ @@ -174,6 +253,7 @@ test("provider smoke rejects a failed model turn with a stable non-secret error" spawnFn: () => child, authProviders: () => new Set(["zai"]), readAuthStore: () => '{"zai":{"type":"api_key","key":"test-only"}}', + readModelsStore: () => undefined, }); let caught: unknown;