Luigit
repositories / pi-ext

pi-ext

bugabingas pi extensions

owned by admin

extensions/masks/index.ts

Raw
import { readFileSync } from "node:fs";
import type {
	ExtensionAPI,
	ExtensionContext,
	InputEvent,
} from "@earendil-works/pi-coding-agent";
import {
	DynamicBorder,
	parseFrontmatter,
} from "@earendil-works/pi-coding-agent";
import {
	Container,
	fuzzyFilter,
	Input,
	SelectList,
	Spacer,
	Text,
} from "@earendil-works/pi-tui";
import { closeDebug, dbg, span } from "./src/debug.ts";
import {
	resolveSetting,
	type SettingDeclaration,
} from "./src/pi-ext-settings.ts";

const DEFAULT_PICKER_SHORTCUT = "alt+m";
type ThinkingLevel = Parameters<ExtensionAPI["setThinkingLevel"]>[0];
const THINKING_LEVELS = [
	"off",
	"minimal",
	"low",
	"medium",
	"high",
	"xhigh",
	"max",
] as const satisfies readonly ThinkingLevel[];

type Shortcut = Parameters<ExtensionAPI["registerShortcut"]>[0];

interface Mask {
	name: string;
	model: string;
	thinkingLevel: ThinkingLevel;
	shortcut?: string;
}

interface MasksConfig {
	pickerShortcut: string;
	items: Mask[];
}

interface ParsedMasks {
	config: MasksConfig;
	errors: string[];
}

function isRecord(value: unknown): value is Record<string, unknown> {
	return typeof value === "object" && value !== null && !Array.isArray(value);
}

function nonEmptyString(value: unknown): string | undefined {
	return typeof value === "string" && value.trim() ? value.trim() : undefined;
}

function isThinkingLevel(value: unknown): value is ThinkingLevel {
	return THINKING_LEVELS.includes(value as ThinkingLevel);
}

function parseMasksSettings(value: unknown): ParsedMasks {
	const errors: string[] = [];
	const config: MasksConfig = {
		pickerShortcut: DEFAULT_PICKER_SHORTCUT,
		items: [],
	};
	if (value === undefined) return { config, errors };
	if (!isRecord(value)) {
		return { config, errors: ["masks must be an object"] };
	}

	if (value.pickerShortcut !== undefined) {
		const shortcut = nonEmptyString(value.pickerShortcut);
		if (shortcut) config.pickerShortcut = shortcut;
		else errors.push("masks.pickerShortcut must be a non-empty string");
	}

	if (value.items === undefined) return { config, errors };
	if (!Array.isArray(value.items)) {
		errors.push("masks.items must be an array");
		return { config, errors };
	}

	const names = new Set<string>();
	const shortcuts = new Set<string>([config.pickerShortcut]);
	for (const [index, valueItem] of value.items.entries()) {
		const path = `masks.items[${index}]`;
		if (!isRecord(valueItem)) {
			errors.push(`${path} must be an object`);
			continue;
		}

		const name = nonEmptyString(valueItem.name);
		const model = nonEmptyString(valueItem.model);
		const thinkingLevel = valueItem.thinkingLevel;
		let valid = true;
		if (!name) {
			errors.push(`${path}.name must be a non-empty string`);
			valid = false;
		} else if (names.has(name)) {
			errors.push(`${path}.name duplicates "${name}"`);
			valid = false;
		}
		if (
			!model ||
			model.startsWith("/") ||
			model.endsWith("/") ||
			!model.includes("/")
		) {
			errors.push(`${path}.model must be provider/model`);
			valid = false;
		}
		if (!isThinkingLevel(thinkingLevel)) {
			errors.push(`${path}.thinkingLevel is invalid`);
			valid = false;
		}
		if (!valid || !name || !model || !isThinkingLevel(thinkingLevel)) continue;

		let shortcut: string | undefined;
		if (valueItem.shortcut !== undefined) {
			shortcut = nonEmptyString(valueItem.shortcut);
			if (!shortcut) {
				errors.push(`${path}.shortcut must be a non-empty string`);
			} else if (shortcuts.has(shortcut)) {
				errors.push(`${path}.shortcut duplicates "${shortcut}"`);
				shortcut = undefined;
			}
		}

		names.add(name);
		if (shortcut) shortcuts.add(shortcut);
		config.items.push({ name, model, thinkingLevel, shortcut });
	}
	return { config, errors };
}

// Parsing keeps valid masks and collects per-field errors, so it never rejects.
const MASKS_SETTING: SettingDeclaration<ParsedMasks> = {
	key: "masks",
	parse: parseMasksSettings,
	default: parseMasksSettings(undefined),
};

// Shortcuts register at load time, before any session or project trust exists,
// so masks stay user-only: project settings are never read.
function readMasksSettings(pi: ExtensionAPI): ParsedMasks {
	const setting = resolveSetting(
		pi,
		{ cwd: process.cwd(), isProjectTrusted: () => false },
		MASKS_SETTING,
	);
	return setting.ok
		? setting.value
		: { config: MASKS_SETTING.default.config, errors: [setting.error] };
}

function splitModelRef(model: string): [provider: string, modelId: string] {
	const separator = model.indexOf("/");
	return [model.slice(0, separator), model.slice(separator + 1)];
}

function fuzzyMasks(items: Mask[], query: string): Mask[] {
	if (!query.trim()) return items;
	return fuzzyFilter(
		items,
		query,
		(mask) =>
			`${mask.name} ${mask.model} ${mask.thinkingLevel} ${mask.shortcut ?? ""}`,
	);
}

async function chooseFuzzyMask(
	items: Mask[],
	ctx: ExtensionContext,
): Promise<Mask | undefined> {
	const selectedName = await ctx.ui.custom<string | undefined>(
		(tui, theme, keybindings, done) => {
			const container = new Container();
			const input = new Input();
			input.focused = true;
			let list: SelectList;

			function rebuild(): void {
				const filtered = fuzzyMasks(items, input.getValue());
				list = new SelectList(
					filtered.map((mask) => ({
						value: mask.name,
						label: mask.name,
						description: `${mask.model}:${mask.thinkingLevel}${mask.shortcut ? ` · ${mask.shortcut}` : ""}`,
					})),
					Math.min(Math.max(filtered.length, 1), 8),
					{
						selectedPrefix: (text) => theme.fg("accent", text),
						selectedText: (text) => theme.fg("accent", text),
						description: (text) => theme.fg("muted", text),
						scrollInfo: (text) => theme.fg("dim", text),
						noMatch: (text) => theme.fg("warning", text),
					},
				);
				list.onSelect = (item) => done(item.value);
				list.onCancel = () => done(undefined);

				container.clear();
				container.addChild(
					new DynamicBorder((text) => theme.fg("accent", text)),
				);
				container.addChild(
					new Text(theme.fg("accent", theme.bold("Masks")), 1, 0),
				);
				container.addChild(new Text(theme.fg("dim", "Fuzzy filter"), 1, 0));
				container.addChild(input);
				container.addChild(new Spacer(1));
				container.addChild(list);
				container.addChild(
					new Text(
						theme.fg(
							"dim",
							"type to filter · ↑↓ select · enter equip · esc cancel",
						),
						1,
						0,
					),
				);
				container.addChild(
					new DynamicBorder((text) => theme.fg("accent", text)),
				);
			}

			rebuild();
			return {
				render: (width: number) => container.render(width),
				invalidate: () => container.invalidate(),
				handleInput(data: string) {
					if (
						keybindings.matches(data, "tui.select.up") ||
						keybindings.matches(data, "tui.select.down")
					) {
						list.handleInput(data);
						tui.requestRender();
						return;
					}
					if (keybindings.matches(data, "tui.select.confirm")) {
						const selected = list.getSelectedItem();
						if (selected) done(selected.value);
						return;
					}
					if (keybindings.matches(data, "tui.select.cancel")) {
						done(undefined);
						return;
					}

					const previous = input.getValue();
					input.handleInput(data);
					if (input.getValue() !== previous) rebuild();
					tui.requestRender();
				},
			};
		},
	);
	return items.find((mask) => mask.name === selectedName);
}

function promptCommandName(text: string): string | undefined {
	return /^\/([^\s]+)(?:\s|$)/u.exec(text)?.[1];
}

function parsePromptTools(
	value: unknown,
	registeredTools: ReadonlySet<string>,
): { tools?: string[]; error?: string } {
	if (!Array.isArray(value) || value.length === 0)
		return { error: "tools must be a non-empty array" };
	const tools: string[] = [];
	const seen = new Set<string>();
	for (const [index, item] of value.entries()) {
		const name = nonEmptyString(item);
		if (!name) return { error: `tools[${index}] must be a non-empty string` };
		if (seen.has(name)) return { error: `tools duplicates "${name}"` };
		if (!registeredTools.has(name)) return { error: `unknown tool "${name}"` };
		seen.add(name);
		tools.push(name);
	}
	return { tools };
}

async function applyMask(
	pi: ExtensionAPI,
	mask: Mask,
	ctx: ExtensionContext,
): Promise<boolean> {
	const end = span?.("mask.apply");
	const [provider, modelId] = splitModelRef(mask.model);
	const model = ctx.modelRegistry.find(provider, modelId);
	if (!model) {
		end?.("error", { type: "model_missing" });
		ctx.ui.notify(
			`Mask "${mask.name}": model ${mask.model} not found`,
			"error",
		);
		return false;
	}
	try {
		if (!(await pi.setModel(model))) {
			end?.("error", { type: "credentials" });
			ctx.ui.notify(
				`Mask "${mask.name}": no credentials for ${mask.model}`,
				"error",
			);
			return false;
		}
		pi.setThinkingLevel(mask.thinkingLevel);
		end?.();
		return true;
	} catch (error) {
		end?.("error", { type: "apply" });
		throw error;
	}
}

export default function masksExtension(pi: ExtensionAPI): void {
	const { config, errors } = readMasksSettings(pi);
	const deferred: Pick<InputEvent, "text" | "images">[] = [];
	let preparing = false;
	let transitionGeneration = 0;
	let manualTransition = false;
	let restoring = false;
	let applyPromise: Promise<boolean> | undefined;
	let restorePromise: Promise<void> | undefined;
	let detachAbort: (() => void) | undefined;
	const terminatingTools = new Set<string>();
	let restoreState:
		| {
				model: NonNullable<ExtensionContext["model"]>;
				thinkingLevel: ThinkingLevel;
				tools?: string[];
				modelPending: boolean;
				toolsPending: boolean;
		  }
		| undefined;

	function showQueue(ctx: ExtensionContext): void {
		ctx.ui.setStatus(
			"masks-queue",
			deferred.length
				? `masks: ${deferred.length} deferred · /masks-cancel`
				: undefined,
		);
	}

	function cancelQueue(ctx: ExtensionContext): void {
		const count = deferred.length;
		deferred.length = 0;
		showQueue(ctx);
		if (count)
			ctx.ui.notify(`Cancelled ${count} deferred masked prompts`, "warning");
	}

	pi.registerCommand("masks-cancel", {
		description:
			"Cancel deferred masked prompts without interrupting the current run",
		handler: async (_args, ctx) => cancelQueue(ctx),
	});

	async function readyForManualMask(ctx: ExtensionContext): Promise<boolean> {
		if (preparing || restoring || applyPromise) {
			ctx.ui.notify("Prompt mask transition in progress", "warning");
			return false;
		}
		if (!restoreState) return true;
		if (!ctx.isIdle()) {
			ctx.ui.notify("Temporary prompt mask is active", "warning");
			return false;
		}
		await restoreTemporaryMask(ctx);
		return restoreState === undefined;
	}

	async function applyManualMask(
		mask: Mask,
		ctx: ExtensionContext,
	): Promise<void> {
		if (manualTransition) {
			ctx.ui.notify("Prompt mask transition in progress", "warning");
			return;
		}
		manualTransition = true;
		let ownedPromise: Promise<boolean> | undefined;
		try {
			if (!(await readyForManualMask(ctx))) return;
			ownedPromise = applyMask(pi, mask, ctx);
			applyPromise = ownedPromise;
			await ownedPromise;
		} finally {
			if (applyPromise === ownedPromise) applyPromise = undefined;
			manualTransition = false;
		}
	}

	async function chooseMask(ctx: ExtensionContext): Promise<void> {
		if (!(await readyForManualMask(ctx))) return;
		if (config.items.length === 0) {
			ctx.ui.notify("No masks configured in settings.json", "warning");
			return;
		}
		let mask: Mask | undefined;
		if (ctx.mode === "tui") {
			mask = await chooseFuzzyMask(config.items, ctx);
		} else {
			const options = config.items.map(
				(item) => `${item.name}  ${item.model}:${item.thinkingLevel}`,
			);
			const selected = await ctx.ui.select("Masks", options);
			const index = selected === undefined ? -1 : options.indexOf(selected);
			mask = index >= 0 ? config.items[index] : undefined;
		}
		if (mask) await applyManualMask(mask, ctx);
	}

	pi.registerShortcut(config.pickerShortcut as Shortcut, {
		description: "Choose model and thinking mask",
		handler: chooseMask,
	});
	for (const mask of config.items) {
		if (!mask.shortcut) continue;
		pi.registerShortcut(mask.shortcut as Shortcut, {
			description: `Equip mask ${mask.name}`,
			handler: async (ctx) => applyManualMask(mask, ctx),
		});
	}

	pi.registerCommand("masks", {
		description: "Switch model and thinking mask",
		handler: async (args, ctx) => {
			const name = args?.trim();
			if (!name) {
				await chooseMask(ctx);
				return;
			}
			const mask = config.items.find((item) => item.name === name);
			if (!mask) {
				ctx.ui.notify(`Unknown mask "${name}"`, "error");
				return;
			}
			await applyManualMask(mask, ctx);
		},
	});

	pi.on("input", async (event, ctx) => {
		if (preparing || manualTransition || restoring || applyPromise) {
			ctx.ui.notify(
				"Masked prompt is transitioning; retry after it starts, or /reload after a dispatch failure",
				"warning",
			);
			return { action: "handled" };
		}
		if (ctx.isIdle() && restoreState) {
			await restoreTemporaryMask(ctx);
			if (restoreState) return { action: "handled" };
		}
		const commandName = promptCommandName(event.text);
		if (!commandName) return { action: "continue" };
		const command = pi
			.getCommands()
			.find((item) => item.source === "prompt" && item.name === commandName);
		if (!command) return { action: "continue" };

		let frontmatter: Record<string, unknown>;
		try {
			frontmatter = parseFrontmatter(
				readFileSync(command.sourceInfo.path, "utf8"),
			).frontmatter;
		} catch (error) {
			ctx.ui.notify(
				`Prompt /${commandName}: ${error instanceof Error ? error.message : String(error)}`,
				"error",
			);
			return { action: "handled" };
		}
		if (!Object.hasOwn(frontmatter, "mask")) return { action: "continue" };
		const maskName = nonEmptyString(frontmatter.mask);
		if (!maskName) {
			ctx.ui.notify(
				`Prompt /${commandName}: mask must be a non-empty string`,
				"error",
			);
			return { action: "handled" };
		}
		const mask = config.items.find((item) => item.name === maskName);
		if (!mask) {
			ctx.ui.notify(
				`Prompt /${commandName}: unknown mask "${maskName}"`,
				"error",
			);
			return { action: "handled" };
		}
		if (!ctx.isIdle()) {
			if (event.streamingBehavior === "steer") {
				ctx.ui.notify(
					`Prompt /${commandName}: use follow-up delivery, not steering`,
					"warning",
				);
				return { action: "handled" };
			}
			deferred.push({ text: event.text, images: event.images });
			showQueue(ctx);
			return { action: "handled" };
		}
		let tools: string[] | undefined;
		if (Object.hasOwn(frontmatter, "tools")) {
			const parsed = parsePromptTools(
				frontmatter.tools,
				new Set(pi.getAllTools().map((tool) => tool.name)),
			);
			if (parsed.error) {
				if (event.source === "extension") cancelQueue(ctx);
				ctx.ui.notify(`Prompt /${commandName}: ${parsed.error}`, "error");
				return { action: "handled" };
			}
			tools = parsed.tools;
		}
		if (!ctx.model) {
			ctx.ui.notify(
				`Mask "${mask.name}": no current model to restore`,
				"error",
			);
			return { action: "handled" };
		}
		restoreState = {
			model: ctx.model,
			thinkingLevel: pi.getThinkingLevel(),
			...(tools ? { tools: [...pi.getActiveTools()] } : {}),
			modelPending: true,
			toolsPending: tools !== undefined,
		};
		const generation = transitionGeneration;
		preparing = true;
		applyPromise = (async () => {
			if (!(await applyMask(pi, mask, ctx))) return false;
			if (!tools) return true;
			pi.setActiveTools(tools);
			const active = pi.getActiveTools();
			if (
				active.length !== tools.length ||
				active.some((name, index) => name !== tools[index])
			)
				throw new Error("tool profile could not be applied exactly");
			return true;
		})();
		let applied = false;
		try {
			applied = await applyPromise;
		} catch (error) {
			ctx.ui.notify(`Failed to equip prompt mask: ${String(error)}`, "error");
		} finally {
			applyPromise = undefined;
		}
		if (!applied || generation !== transitionGeneration) {
			if (!applied) cancelQueue(ctx);
			await restoreTemporaryMask(ctx);
			preparing = false;
			return { action: "handled" };
		}
		return { action: "continue" };
	});

	pi.on("agent_start", (_event, ctx) => {
		preparing = false;
		detachAbort?.();
		const signal = ctx.signal;
		const cancel = () => cancelQueue(ctx);
		signal?.addEventListener("abort", cancel, { once: true });
		detachAbort = () => signal?.removeEventListener("abort", cancel);
		if (signal?.aborted) cancel();
	});

	async function restoreTemporaryMask(ctx: ExtensionContext): Promise<void> {
		if (restorePromise) return restorePromise;
		const previous = restoreState;
		if (!previous) return;
		restoring = true;
		restorePromise = (async () => {
			const failures: string[] = [];
			if (previous.modelPending) {
				try {
					if (!(await pi.setModel(previous.model)))
						throw new Error("model unavailable");
					pi.setThinkingLevel(previous.thinkingLevel);
					previous.modelPending = false;
				} catch (error) {
					dbg?.("restore.failure", { type: "model" });
					failures.push(`model: ${String(error)}`);
				}
			}
			if (previous.toolsPending && previous.tools) {
				try {
					pi.setActiveTools(previous.tools);
					const active = pi.getActiveTools();
					if (
						active.length !== previous.tools.length ||
						active.some((name, index) => name !== previous.tools?.[index])
					)
						throw new Error("tool snapshot could not be restored exactly");
					previous.toolsPending = false;
				} catch (error) {
					dbg?.("restore.failure", { type: "tools" });
					failures.push(`tools: ${String(error)}`);
				}
			}
			if (!previous.modelPending && !previous.toolsPending) {
				if (restoreState === previous) restoreState = undefined;
				return;
			}
			cancelQueue(ctx);
			ctx.ui.notify(
				`Failed to restore prompt mask: ${failures.join("; ")}`,
				"error",
			);
			ctx.abort();
		})();
		try {
			await restorePromise;
		} finally {
			restorePromise = undefined;
			restoring = false;
		}
	}

	pi.on("turn_start", () => terminatingTools.clear());
	pi.on("tool_execution_end", (event) => {
		if (event.result.terminate === true) terminatingTools.add(event.toolCallId);
	});
	pi.on("turn_end", async (event, ctx) => {
		if (event.message.role !== "assistant") return;
		const tools = event.message.content.filter(
			(part) => part.type === "toolCall",
		);
		if (
			tools.length
				? tools.every((tool) => terminatingTools.has(tool.id))
				: event.message.stopReason === "stop"
		) {
			await restoreTemporaryMask(ctx);
		}
	});
	pi.on("agent_end", async (event, ctx) => {
		if (
			event.messages.some(
				(message) =>
					message.role === "assistant" &&
					(message.stopReason === "aborted" || message.stopReason === "error"),
			)
		)
			cancelQueue(ctx);
	});
	pi.on("agent_settled", async (_event, ctx) => {
		if (!ctx.isIdle()) return;
		await restoreTemporaryMask(ctx);
		detachAbort?.();
		detachAbort = undefined;
		if (restoreState || preparing || !ctx.isIdle()) return;
		const next = deferred.shift();
		if (!next) return;
		showQueue(ctx);
		pi.sendUserMessage(
			[{ type: "text", text: next.text }, ...(next.images ?? [])],
			{ expandPromptTemplates: true, deliverAs: "followUp" },
		);
	});
	pi.on("session_start", (_event, ctx) => {
		dbg?.("session.start", { mode: ctx.mode });
		for (const error of errors) ctx.ui.notify(error, "warning");
	});
	async function reset(ctx: ExtensionContext): Promise<void> {
		transitionGeneration++;
		preparing = false;
		detachAbort?.();
		detachAbort = undefined;
		cancelQueue(ctx);
		await applyPromise;
		await restoreTemporaryMask(ctx);
	}
	pi.on("session_tree", async (_event, ctx) => reset(ctx));
	pi.on("session_shutdown", async (_event, ctx) => {
		dbg?.("session.shutdown", { mode: ctx.mode });
		try {
			await reset(ctx);
		} finally {
			closeDebug();
		}
	});
}

export const __test = {
	applyMask,
	chooseFuzzyMask,
	fuzzyMasks,
	parseMasksSettings,
	parsePromptTools,
	promptCommandName,
	splitModelRef,
};