Luigit
repositories / pi-ext

pi-ext

bugabingas pi extensions

owned by admin

extensions/rtk/index.ts

Raw
import type {
	ExecOptions,
	ExecResult,
	ExtensionAPI,
	ExtensionContext,
} from "@earendil-works/pi-coding-agent";
import { isToolCallEventType } from "@earendil-works/pi-coding-agent";
import { closeDebug, dbg, span } from "./src/debug.ts";
import { findExecutable } from "./src/pi-ext-executable.ts";
import {
	MAX_TIMER_MS,
	parseNumberSetting,
	resolveSetting,
	type SettingDeclaration,
} from "./src/pi-ext-settings.ts";
import { isStaleContextError } from "./src/pi-ext-stale-context.ts";

const TIMEOUT_SETTING: SettingDeclaration<number> = {
	key: "rtk.rewriteTimeoutMs",
	env: "PI_RTK_REWRITE_TIMEOUT_MS",
	default: 2_000,
	parse: (raw, source) => {
		const value = parseNumberSetting(raw, source);
		if (
			value === undefined ||
			!Number.isInteger(value) ||
			value <= 0 ||
			value > MAX_TIMER_MS
		)
			throw new Error(`expected an integer from 1 to ${MAX_TIMER_MS}`);
		return value;
	},
};
const COMMAND_PREFIX_END = "# pi-command-prefix-end\n";
const MIN_MAJOR = 0;
const MIN_MINOR = 23;

type VersionTuple = [number, number, number];

type RtkStatus =
	| { ok: true }
	| { ok: false; reason: "missing" | "old"; version?: string };

export function parseRtkVersion(text: string): VersionTuple | null {
	const match = text.match(/(\d+)\.(\d+)\.(\d+)/);
	if (!match) return null;
	return [Number(match[1]), Number(match[2]), Number(match[3])];
}

function supported(version: VersionTuple | null): boolean {
	if (version === null) return false;
	const [major, minor] = version;
	return major > MIN_MAJOR || (major === MIN_MAJOR && minor >= MIN_MINOR);
}

export function shouldSkipCommand(command: string): boolean {
	const trimmed = command.trimStart();
	if (trimmed.trim() === "") return true;
	if (trimmed === "rtk" || trimmed.startsWith("rtk ")) return true;
	return /^(?:env\s+)?RTK_DISABLED=1(?:\s|$)/u.test(trimmed);
}

export function splitCommandPrefix(command: string): {
	prefix: string;
	command: string;
} {
	const end = command.indexOf(COMMAND_PREFIX_END);
	if (end < 0) return { prefix: "", command };
	const bodyStart = end + COMMAND_PREFIX_END.length;
	return {
		prefix: command.slice(0, bodyStart),
		command: command.slice(bodyStart),
	};
}

function hasShellOperators(command: string): boolean {
	return /[|;<>`]|\$\(|&&|\|\|/u.test(command);
}

export function localFallbackRewrite(command: string): string | undefined {
	const trimmed = command.trim();
	if (hasShellOperators(trimmed)) return undefined;
	if (/^(journalctl|ps)(\s|$)/u.test(trimmed)) return `rtk summary ${trimmed}`;
	if (/^systemctl\s+(?:--user\s+)?status(\s|$)/u.test(trimmed))
		return `rtk summary ${trimmed}`;
	if (
		/^hyprctl\s+(clients|monitors|workspaces|activewindow|activeworkspace)(\s|$)/u.test(
			trimmed,
		)
	)
		return `rtk summary ${trimmed}`;
	return undefined;
}

type RewriteOutcome =
	| { kind: "rewrite"; command: string }
	| { kind: "fallback"; reason: "no_rewrite" }
	| {
			kind: "unchanged";
			reason: "no_rewrite" | "killed" | "unexpected_exit" | "thrown";
	  };

async function rtkCommand(): Promise<string> {
	if (process.platform !== "win32") return "rtk";
	const exts = (process.env.PATHEXT || ".COM;.EXE;.BAT;.CMD")
		.split(";")
		.filter(Boolean);
	return findExecutable(exts.map((ext) => `rtk${ext.toLowerCase()}`)) ?? "rtk";
}

async function execRtk(
	pi: ExtensionAPI,
	args: string[],
	options?: ExecOptions,
): Promise<ExecResult> {
	const command = await rtkCommand();
	if (process.platform === "win32" && /\.(?:cmd|bat)$/iu.test(command)) {
		return await pi.exec(
			process.env.ComSpec || process.env.COMSPEC || "cmd.exe",
			["/d", "/c", "call", command, ...args],
			options,
		);
	}
	return await pi.exec(command, args, options);
}

function execCode(error: unknown): number | null {
	if (typeof error !== "object" || error === null) return null;
	const maybe = error as {
		code?: unknown;
		exitCode?: unknown;
		status?: unknown;
	};
	for (const value of [maybe.code, maybe.exitCode, maybe.status])
		if (typeof value === "number") return value;
	return null;
}

async function rewriteCommand(
	pi: ExtensionAPI,
	command: string,
	timeout: number,
	signal?: AbortSignal,
): Promise<RewriteOutcome> {
	try {
		const result = await execRtk(pi, ["rewrite", command], {
			timeout,
			signal,
		});
		if (result.killed) return { kind: "unchanged", reason: "killed" };
		if (result.code === 1) return { kind: "fallback", reason: "no_rewrite" };
		if (result.code !== 0 && result.code !== 3)
			return { kind: "unchanged", reason: "unexpected_exit" };
		const rewritten = result.stdout?.trim();
		if (!rewritten || rewritten === command)
			return { kind: "fallback", reason: "no_rewrite" };
		return { kind: "rewrite", command: rewritten };
	} catch (error) {
		return execCode(error) === 1
			? { kind: "fallback", reason: "no_rewrite" }
			: { kind: "unchanged", reason: "thrown" };
	}
}

async function rtkStatus(
	pi: ExtensionAPI,
	timeout: number,
): Promise<RtkStatus> {
	const finish = span?.("rtk.status");
	let outcome: "available" | "missing" | "unsupported" = "missing";
	try {
		const result = await execRtk(pi, ["--version"], { timeout });
		if (result.code !== 0) return { ok: false, reason: "missing" };
		const text = (result.stdout ?? result.stderr ?? "").trim();
		if (supported(parseRtkVersion(text))) {
			outcome = "available";
			return { ok: true };
		}
		outcome = "unsupported";
		return { ok: false, reason: "old", version: text };
	} catch {
		return { ok: false, reason: "missing" };
	} finally {
		finish?.("finish", { outcome });
	}
}

function notifyUser(
	pi: ExtensionAPI,
	ctx: ExtensionContext | undefined,
	message: string,
): boolean {
	try {
		const notify = ctx?.hasUI
			? ctx.ui.notify.bind(ctx.ui)
			: (
					pi as {
						ui?: { notify?: (message: string, level: "warning") => void };
					}
				).ui?.notify;
		if (!notify) return false;
		notify(message, "warning");
		return true;
	} catch (error) {
		if (isStaleContextError(error)) return false;
		throw error;
	}
}

function warnUnavailable(
	status: RtkStatus,
	pi: ExtensionAPI,
	ctx?: ExtensionContext,
): boolean {
	if (status.ok) return false;
	const suffix =
		status.reason === "old"
			? `too old (${status.version || "unknown"}); need >= 0.23.0`
			: "not found on PATH";
	return notifyUser(pi, ctx, `rtk disabled: ${suffix}`);
}

export const __test = {
	localFallbackRewrite,
	parseRtkVersion,
	shouldSkipCommand,
	splitCommandPrefix,
};

export default function rtk(pi: ExtensionAPI) {
	let generation = 0;
	let announced = false;
	let rewriteCache = new Map<string, RewriteOutcome>();
	let timeout = TIMEOUT_SETTING.default;
	let initialization:
		| {
				generation: number;
				promise: Promise<RtkStatus>;
		  }
		| undefined;

	const ensureInitialized = async (ctx: ExtensionContext) => {
		if (initialization?.generation !== generation) {
			const ticket = generation;
			const promise = rtkStatus(pi, timeout);
			initialization = { generation: ticket, promise };
		}

		const ticket = initialization.generation;
		const support = await initialization.promise;
		if (ticket !== generation) return undefined;
		if (!announced) {
			announced = support.ok || warnUnavailable(support, pi, ctx);
		}
		return support;
	};

	const reset = () => {
		generation++;
		announced = false;
		initialization = undefined;
		rewriteCache = new Map();
	};
	pi.on("session_start", (_event, ctx) => {
		dbg?.("session.start");
		reset();
		const setting = resolveSetting(pi, ctx, TIMEOUT_SETTING);
		timeout = setting.ok ? setting.value : TIMEOUT_SETTING.default;
		if (!setting.ok)
			notifyUser(
				pi,
				ctx,
				`${setting.error}; using ${TIMEOUT_SETTING.default}ms`,
			);
	});
	pi.on("session_shutdown", () => {
		dbg?.("session.shutdown");
		reset();
		closeDebug();
	});
	pi.on("tool_call", async (event, ctx) => {
		if (!isToolCallEventType("bash", event)) return;
		const support = await ensureInitialized(ctx);
		if (!support?.ok) return;
		const framed = splitCommandPrefix(event.input.command);
		if (shouldSkipCommand(framed.command)) {
			dbg?.("rtk.rewrite", { outcome: "skipped" });
			return;
		}
		const cache = rewriteCache;
		const cached = cache.has(framed.command);
		let outcome = cache.get(framed.command);
		if (!outcome) {
			outcome = await rewriteCommand(pi, framed.command, timeout, ctx.signal);
			if (outcome.kind === "rewrite") cache.set(framed.command, outcome);
		}
		const rewritten =
			outcome.kind === "rewrite"
				? outcome.command
				: outcome.kind === "fallback"
					? localFallbackRewrite(framed.command)
					: undefined;
		if (rewritten) event.input.command = framed.prefix + rewritten;
		const result: "rewritten" | "cached_rewrite" | "fallback" | "unchanged" =
			outcome.kind === "rewrite"
				? cached
					? "cached_rewrite"
					: "rewritten"
				: rewritten
					? "fallback"
					: "unchanged";
		dbg?.("rtk.rewrite", {
			outcome: result,
			...(outcome.kind === "rewrite" ? {} : { reason: outcome.reason }),
		});
	});
}