Luigit
repositories / pi-ext

pi-ext

bugabingas pi extensions

owned by admin

extensions/rtk/__tests__/rtk.test.ts

Raw
import { mkdtempSync, rmSync } from "node:fs";
import { tmpdir } from "node:os";
import path from "node:path";
import { afterAll, beforeAll, describe, expect, test, vi } from "vitest";
import { sandboxEnv } from "../../../test/harness";
import rtkExtension, { __test } from "../index.ts";

// session_start resolves settings; keep the developer's real user settings out of these tests.
const originalEnv = { ...process.env };
const replaceEnv = (env: NodeJS.ProcessEnv) => {
	for (const key of Object.keys(process.env)) delete process.env[key];
	Object.assign(process.env, env);
};
let sandbox: string;
beforeAll(() => {
	sandbox = mkdtempSync(path.join(tmpdir(), "pi-ext-rtk-unit-"));
	replaceEnv(sandboxEnv(sandbox));
});
afterAll(() => {
	replaceEnv(originalEnv);
	rmSync(sandbox, { recursive: true, force: true });
});

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

type Handler = (event: any, ctx: any) => unknown | Promise<unknown>;

function fakePi(exec: (command: string, args: string[]) => Promise<any>) {
	const handlers: Record<string, Handler[]> = {};
	const notes: Array<[string, string]> = [];
	return {
		handlers,
		notes,
		ctx: {
			cwd: process.cwd(),
			isProjectTrusted: () => false,
			hasUI: true,
			ui: { notify: vi.fn((message, level) => notes.push([message, level])) },
		},
		pi: {
			exec: vi.fn(exec),
			on: vi.fn((event: string, handler: Handler) => {
				if (!handlers[event]) handlers[event] = [];
				handlers[event].push(handler);
			}),
			ui: { notify: vi.fn((message, level) => notes.push([message, level])) },
		},
	};
}

async function loadWithRewrite(code: number, stdout = "") {
	const fx = fakePi(async (_command, args) =>
		args[0] === "--version"
			? { code: 0, stdout: "rtk 0.42.4" }
			: { code, stdout },
	);
	rtkExtension(fx.pi as never);
	return fx;
}

async function emit(
	fx: Awaited<ReturnType<typeof loadWithRewrite>>,
	command: string,
	toolName = "bash",
) {
	const event = { toolName, input: { command } };
	await fx.handlers.tool_call?.[0]?.(event, {
		...fx.ctx,
		signal: AbortSignal.abort(),
	});
	return event.input.command;
}

describe("rtk extension", () => {
	test("parses versions and skip guards", () => {
		expect(parseRtkVersion("rtk 0.42.4")).toEqual([0, 42, 4]);
		expect(parseRtkVersion("garbage")).toBeNull();
		expect(shouldSkipCommand("rtk git status")).toBe(true);
		expect(shouldSkipCommand("env RTK_DISABLED=1 git status")).toBe(true);
		expect(shouldSkipCommand("git status")).toBe(false);
	});

	test("preserves registered command prefixes around rewrites", async () => {
		const prefix = "# policy\n# pi-command-prefix-end\n";
		expect(splitCommandPrefix(`${prefix}git status`)).toEqual({
			prefix,
			command: "git status",
		});
		expect(
			await emit(
				await loadWithRewrite(0, "rtk git status\n"),
				`${prefix}git status`,
			),
		).toBe(`${prefix}rtk git status`);
	});

	test("maps local noisy-output misses only", () => {
		expect(localFallbackRewrite("journalctl --user -n 100 --no-pager")).toBe(
			"rtk summary journalctl --user -n 100 --no-pager",
		);
		expect(localFallbackRewrite("systemctl --user status pipewire")).toBe(
			"rtk summary systemctl --user status pipewire",
		);
		expect(localFallbackRewrite("hyprctl clients")).toBe(
			"rtk summary hyprctl clients",
		);
		expect(localFallbackRewrite("ps aux")).toBe("rtk summary ps aux");
		expect(
			localFallbackRewrite("mise run test -- extensions/rtk"),
		).toBeUndefined();
		expect(
			localFallbackRewrite("mise exec -- biome check extensions/rtk"),
		).toBeUndefined();
		expect(
			localFallbackRewrite("journalctl --user | grep pipewire"),
		).toBeUndefined();
	});

	test("warns once and leaves commands unchanged when rtk is missing", async () => {
		const fx = fakePi(async () => ({ code: 127, stdout: "" }));
		rtkExtension(fx.pi as never);
		expect(fx.handlers.tool_call).toHaveLength(1);
		expect(fx.notes).toEqual([]);
		expect(await emit(fx as never, "ps aux")).toBe("ps aux");
		expect(fx.notes).toEqual([["rtk disabled: not found on PATH", "warning"]]);
	});

	test("registers without probing during extension load", () => {
		const fx = fakePi(async () => ({ code: 127, stdout: "" }));
		rtkExtension(fx.pi as never);
		expect(fx.pi.exec).not.toHaveBeenCalled();
		expect(fx.notes).toEqual([]);
	});

	test("defers the availability probe until the first bash call", async () => {
		const fx = fakePi(async () => ({ code: 127, stdout: "" }));
		rtkExtension(fx.pi as never);
		const result = fx.handlers.session_start?.[0]?.({}, fx.ctx);
		expect(result).toBeUndefined();
		await new Promise((resolve) => setTimeout(resolve, 0));
		expect(fx.pi.exec).not.toHaveBeenCalled();
		expect(await emit(fx as never, "ps aux")).toBe("ps aux");
		expect(fx.notes).toEqual([["rtk disabled: not found on PATH", "warning"]]);
	});

	test("warns once and leaves commands unchanged when rtk is old", async () => {
		const fx = fakePi(async () => ({ code: 0, stdout: "rtk 0.22.9" }));
		rtkExtension(fx.pi as never);
		expect(await emit(fx as never, "ps aux")).toBe("ps aux");
		expect(fx.notes).toEqual([
			["rtk disabled: too old (rtk 0.22.9); need >= 0.23.0", "warning"],
		]);
	});

	test("registers only tool_call when rtk is available", async () => {
		const fx = await loadWithRewrite(1);
		expect(fx.pi.on).toHaveBeenCalledWith("tool_call", expect.any(Function));
		await emit(fx, "git status");
		expect(fx.notes).toEqual([]);
	});

	test("rewrites via rtk and falls back only on no-match", async () => {
		expect(
			await emit(await loadWithRewrite(0, "rtk git status\n"), "git status"),
		).toBe("rtk git status");
		expect(
			await emit(await loadWithRewrite(3, "rtk cargo test\n"), "cargo test"),
		).toBe("rtk cargo test");
		expect(await emit(await loadWithRewrite(1), "ps aux")).toBe(
			"rtk summary ps aux",
		);
		expect(await emit(await loadWithRewrite(2), "ps aux")).toBe("ps aux");
	});

	test("uses local fallback when pi.exec throws exit 1", async () => {
		const fx = fakePi(async (_command, args) => {
			if (args[0] === "--version") return { code: 0, stdout: "rtk 0.42.4" };
			throw Object.assign(new Error("no match"), { exitCode: 1 });
		});
		await rtkExtension(fx.pi as never, fx.ctx as never);
		expect(await emit(fx as never, "ps aux")).toBe("rtk summary ps aux");
	});

	test("rewrites find commands without a compatibility probe", async () => {
		const fx = fakePi(async (_command, args) => {
			if (args[0] === "--version") return { code: 0, stdout: "rtk 0.42.4" };
			return { code: 3, stdout: "rtk find . -name AGENTS.md" };
		});
		await rtkExtension(fx.pi as never, fx.ctx as never);
		expect(await emit(fx as never, "find . -name AGENTS.md")).toBe(
			"rtk find . -name AGENTS.md",
		);
		expect(fx.notes).toEqual([]);
	});

	test("caches exact rewrites for the current session", async () => {
		const fx = await loadWithRewrite(0, "rtk git status\n");
		expect(await emit(fx, "git status")).toBe("rtk git status");
		expect(await emit(fx, "git status")).toBe("rtk git status");
		expect(fx.pi.exec).toHaveBeenCalledTimes(2);

		fx.handlers.session_start?.[0]?.({}, fx.ctx);
		expect(await emit(fx, "git status")).toBe("rtk git status");
		expect(fx.pi.exec).toHaveBeenCalledTimes(4);
	});

	test("leaves non-bash and skipped commands untouched", async () => {
		const fx = await loadWithRewrite(0, "rtk git status\n");
		expect(await emit(fx, "git status", "read")).toBe("git status");
		expect(await emit(fx, "rtk git status")).toBe("rtk git status");
	});
});