Luigit
repositories / pi-ext

pi-ext

bugabingas pi extensions

owned by admin

extensions/ultra/__tests__/flags.test.ts

Raw
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<string, FlagValue>;
		flags?: Record<string, FlagValue>;
		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<string, FlagValue>([
				["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<string, FlagValue>([["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<string, FlagValue>([["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, FlagValue>, string][])(
		"rejects invalid step flags %o",
		async (flags, message) => {
			await expect(child({ flags })).rejects.toThrow(message);
		},
	);
});