import { mkdir, mkdtemp, rm, writeFile } from "node:fs/promises"; import { tmpdir } from "node:os"; import { join } from "node:path"; import { createAgentSession, DefaultResourceLoader, SessionManager, } from "@earendil-works/pi-coding-agent"; import { afterEach, describe, expect, it } from "vitest"; import { sandboxEnv, withProcessEnv } from "../../../test/harness"; import { type FlagValue, type HeadlessSession, loadSubagentFlagCatalog, makeSessionFactory, type SessionFactoryDeps, subagentFlagPrompt, } from "../drivers.ts"; // A fixture that behaves like chrome-cdp, firefox-bidi, and nushell: // its tool is registered but switched off at session_start unless its flag is set. const GATED_EXTENSION = ` import { Type } from "typebox"; export default function gated(pi) { pi.registerFlag("gated", { description: "Enable gated_tool at startup", type: "boolean" }); pi.registerFlag("gated-mode", { description: "Mode label", type: "string" }); pi.registerTool({ name: "gated_tool", label: "gated", description: "Reports the mode flag", parameters: Type.Object({}), async execute() { return { content: [{ type: "text", text: String(pi.getFlag("gated-mode")) }], details: undefined }; }, }); pi.on("session_start", () => { const active = pi.getActiveTools().filter((name) => name !== "gated_tool"); pi.setActiveTools(pi.getFlag("gated") === true ? [...active, "gated_tool"] : active); }); } `; describe("sub-agent extension flags", () => { let root: string | undefined; const sessions: HeadlessSession[] = []; afterEach(async () => { for (const session of sessions.splice(0)) session.dispose(); if (root) await rm(root, { recursive: true, force: true }); root = undefined; }); async function child(options: { inheritedFlags?: ReadonlyMap; flags?: Record; tools?: string[]; }) { root = await mkdtemp(join(tmpdir(), "ultra-flags-")); const agentDir = join(root, "agent"); const extension = join(root, "gated", "index.ts"); await mkdir(join(root, "gated"), { recursive: true }); await mkdir(agentDir, { recursive: true }); await writeFile(extension, GATED_EXTENSION); const env = sandboxEnv(root, { env: { PI_CODING_AGENT_DIR: agentDir } }); const cwd = root; return withProcessEnv(env, async () => { const factory = makeSessionFactory({ createAgentSession: createAgentSession as unknown as SessionFactoryDeps["createAgentSession"], DefaultResourceLoader: DefaultResourceLoader as unknown as SessionFactoryDeps["DefaultResourceLoader"], getAgentDir: () => agentDir, cwd, extensionPaths: [extension], inheritedFlags: options.inheritedFlags, }); const session = (await factory({ model: undefined, tools: options.tools ?? ["read", "gated_tool"], customTools: [], sessionManager: SessionManager.inMemory(cwd), ...(options.flags ? { flags: options.flags } : {}), })) as HeadlessSession; sessions.push(session); return session; }); } it("catalogs only extensions that own both tools and flags", async () => { root = await mkdtemp(join(tmpdir(), "ultra-flag-catalog-")); const files = { gated: GATED_EXTENSION, "flag-only": 'export default (pi) => pi.registerFlag("lonely", { type: "boolean" });\n', "tool-only": `import { Type } from "typebox";\nexport default (pi) => pi.registerTool({ name: "plain", label: "plain", description: "plain", parameters: Type.Object({}), async execute() { return { content: [], details: undefined }; } });\n`, }; const paths: string[] = []; for (const [name, source] of Object.entries(files)) { await mkdir(join(root, name), { recursive: true }); await writeFile(join(root, name, "index.ts"), source); paths.push(join(root, name, "index.ts")); } const agentDir = join(root, "agent"); const cwd = root; const catalog = await withProcessEnv( sandboxEnv(root, { env: { PI_CODING_AGENT_DIR: agentDir } }), () => loadSubagentFlagCatalog({ DefaultResourceLoader: DefaultResourceLoader as unknown as SessionFactoryDeps["DefaultResourceLoader"], getAgentDir: () => agentDir, cwd, extensionPaths: paths, }), ); expect(catalog.map((entry) => entry.extension)).toEqual(["gated"]); expect(subagentFlagPrompt(catalog).split("\n").at(-1)).toBe( "- gated: gated_tool; --gated (boolean): Enable gated_tool at startup; --gated-mode (string): Mode label", ); }); it("fails a step whose requested tool its extension switched off, naming the flag", async () => { await expect(child({})).rejects.toThrow( "ultra: requested tools were switched off by their extensions at startup: gated_tool (extension gated; flags: --gated (Enable gated_tool at startup), --gated-mode (Mode label)). Enable them through the owning extension's flag in step.flags.", ); }); it("inherits the parent's CLI flags as Pi parses them", async () => { // `pi --gated some prompt` makes the untyped parser capture a value for the boolean flag. const session = await child({ inheritedFlags: new Map([ ["gated", "some prompt"], ["not-registered-here", true], ]), }); expect(session.getActiveToolNames()).toContain("gated_tool"); }); it("lets step flags enable tools and override inherited values", async () => { const session = await child({ inheritedFlags: new Map([["gated-mode", "parent"]]), flags: { gated: true, "gated-mode": "step" }, }); expect(session.getActiveToolNames()).toContain("gated_tool"); const tool = ( session as unknown as { extensionRunner: { getToolDefinition(name: string): { execute( ...args: unknown[] ): Promise<{ content: { text: string }[] }>; }; }; } ).extensionRunner.getToolDefinition("gated_tool"); const result = await tool.execute("call", {}, undefined, undefined, {}); expect(result.content[0]?.text).toBe("step"); }); it("lets a step force a boolean flag off", async () => { await expect( child({ inheritedFlags: new Map([["gated", true]]), flags: { gated: false }, }), ).rejects.toThrow("switched off by their extensions"); }); it.each([ [ { missing: true }, "ultra: step flag --missing is not registered by any sub-agent extension.", ], [{ gated: "yes" }, "ultra: step flag --gated expects a boolean value."], [ { "gated-mode": true }, "ultra: step flag --gated-mode expects a string value.", ], ] satisfies [Record, string][])( "rejects invalid step flags %o", async (flags, message) => { await expect(child({ flags })).rejects.toThrow(message); }, ); });