Luigit
repositories / pi-ext

pi-ext

bugabingas pi extensions

owned by admin

scripts/context-startup-session.mjs

Raw
import { constants, existsSync } from "node:fs";
import {
	access,
	readFile,
	realpath,
	rename,
	stat,
	writeFile,
} from "node:fs/promises";
import { homedir } from "node:os";
import {
	basename,
	delimiter,
	dirname,
	extname,
	join,
	resolve,
	sep,
} from "node:path";
import { pathToFileURL } from "node:url";
import { PROMPT, snapshot, summarize } from "./context-startup.mjs";

// Resolve the installed executable, not this repository's SDK dependency.
export async function installedPiDist() {
	for (const directory of (process.env.PATH ?? "").split(delimiter)) {
		const candidate = resolve(directory, "pi");
		try {
			await access(candidate, constants.X_OK);
			if (!(await stat(candidate)).isFile()) continue;
		} catch {
			continue;
		}
		let directoryPath = dirname(await realpath(candidate));
		while (true) {
			try {
				const manifest = JSON.parse(
					await readFile(join(directoryPath, "package.json"), "utf8"),
				);
				if (manifest.name === "@earendil-works/pi-coding-agent")
					return join(directoryPath, "dist");
			} catch (error) {
				if (error.code !== "ENOENT")
					throw new Error("Could not locate installed Pi");
			}
			const parent = dirname(directoryPath);
			if (parent === directoryPath) break;
			directoryPath = parent;
		}
		break;
	}
	throw new Error("Could not locate installed Pi");
}

export async function estimateSourceTokens(
	agentDir,
	sdk,
	{ skills, diagnostics },
) {
	const append = await readFile(join(agentDir, "APPEND_SYSTEM.md"), "utf8");
	if (diagnostics.length)
		throw new Error(
			"Skill catalog has diagnostics; refusing a partial estimate",
		);
	const tokens = (content) =>
		sdk.estimateTokens({ role: "user", content, timestamp: 0 });
	const groups = { piAgent: [], extensions: [] };
	let ignoredProjectSkills = 0;
	for (const skill of skills.filter((skill) => !skill.disableModelInvocation)) {
		const source = skill.sourceInfo;
		if (
			source?.origin === "package" ||
			source?.source?.startsWith("extension:")
		) {
			groups.extensions.push(skill);
		} else if (
			resolve(skill.filePath).startsWith(`${resolve(agentDir, "skills")}${sep}`)
		) {
			groups.piAgent.push(skill);
		} else if (source?.scope === "project") {
			ignoredProjectSkills++;
		} else {
			throw new Error("Unattributed skill source; refusing a partial estimate");
		}
	}
	const parts = (items) => {
		const text = sdk.formatSkillsForPrompt(items);
		if (!text) return { shared: "", entries: "" };
		const start = text.indexOf("  <skill>");
		const end = text.lastIndexOf("</available_skills>");
		if (start < 0 || end <= start)
			throw new Error("Pi skill catalog format changed");
		return {
			shared: text.slice(0, start) + text.slice(end),
			entries: text.slice(start, end),
		};
	};
	const piAgent = parts(groups.piAgent);
	const extensions = parts(groups.extensions);
	const shared = piAgent.shared || extensions.shared;
	const appendix = tokens(append);
	const catalog = tokens(shared + piAgent.entries + extensions.entries);
	return {
		appendix,
		descriptions: groups.piAgent
			.map(({ name, description }) => ({
				name,
				description,
				tokens: tokens(description),
			}))
			.sort(
				(a, b) =>
					b.description.length - a.description.length ||
					a.name.localeCompare(b.name),
			),
		piAgent: { tokens: tokens(piAgent.entries), count: groups.piAgent.length },
		extensions: {
			tokens: tokens(extensions.entries),
			count: groups.extensions.length,
		},
		shared: tokens(shared),
		ignoredProjectSkills,
		catalog,
		total: appendix + catalog,
	};
}

export function inspectExtensionRegistrations(
	{ extensions, runtime },
	previous = [],
) {
	const providers = [
		...(runtime.pendingProviderRegistrations ?? []).map(
			({ name, extensionPath }) => ({ name, path: extensionPath }),
		),
		...(runtime.pendingNativeProviderRegistrations ?? []).map(
			({ provider, extensionPath }) => ({
				name: provider.id,
				path: extensionPath,
			}),
		),
	];
	return extensions
		.filter((extension) => !extension.path.startsWith("<inline:"))
		.map((extension) => {
			const path = extension.path;
			const file = basename(path, extname(path));
			return {
				path,
				name: file === "index" ? basename(dirname(path)) : file,
				tools: [...extension.tools.keys()].sort(),
				hooks: [...extension.handlers.keys()].sort(),
				providers: [
					...new Set([
						...(previous.find((entry) => entry.path === path)?.providers ?? []),
						...providers
							.filter((entry) => entry.path === path)
							.map((entry) => entry.name),
					]),
				].sort(),
			};
		});
}

export function resourceOptions(
	kind,
	extensionFactories,
	{ cwd, agentDir, projectTrusted, extensionPaths } = {},
) {
	if (kind === "headless+1msg+no_pi_config") {
		// Keep the real settings and extension discovery. Remove only agent-dir text resources.
		const root = resolve(agentDir);
		const outside = (path, directory) =>
			!resolve(path).startsWith(`${join(root, directory)}${sep}`);
		const projectPrompt = (name) => {
			const path = join(cwd, ".pi", name);
			return projectTrusted && existsSync(path) ? path : "";
		};
		const append = projectPrompt("APPEND_SYSTEM.md");
		return {
			extensionFactories,
			agentsFilesOverride: (base) => ({
				...base,
				agentsFiles: base.agentsFiles.filter(
					(file) => dirname(resolve(file.path)) !== root,
				),
			}),
			skillsOverride: (base) => ({
				...base,
				skills: base.skills.filter((skill) =>
					outside(skill.filePath, "skills"),
				),
			}),
			promptsOverride: (base) => ({
				...base,
				prompts: base.prompts.filter((prompt) =>
					outside(prompt.filePath, "prompts"),
				),
			}),
			systemPrompt: projectPrompt("SYSTEM.md"),
			appendSystemPrompt: append ? [append] : [],
		};
	}
	return kind === "baseline+1msg"
		? {
				noExtensions: true,
				noSkills: true,
				noPromptTemplates: true,
				noThemes: true,
				noContextFiles: true,
				systemPrompt: "",
				appendSystemPrompt: [],
				extensionFactories,
			}
		: {
				extensionFactories,
				...(extensionPaths === undefined
					? {}
					: { noExtensions: true, additionalExtensionPaths: extensionPaths }),
			};
}

export function observeRequests(agent, observe) {
	const stream = agent.streamFunction;
	if (typeof stream !== "function")
		throw new Error("Pi stream API unavailable");
	agent.streamFunction = (model, context, options) => {
		observe(context);
		return stream(model, context, options);
	};
}

// Only counts and public run metadata cross the process boundary.
export async function publish(path, value) {
	await writeFile(`${path}.tmp`, JSON.stringify(value), { mode: 0o600 });
	await rename(`${path}.tmp`, path);
}

export async function runSession(options) {
	const { kind, cwd, agentDir, sourceAgentDir, resultPath, selection } =
		options;
	const baseline = kind === "baseline+1msg";
	process.env.PI_CODING_AGENT_DIR = agentDir;
	process.env.PI_OFFLINE = "1";
	process.env.AI_AGENT = "pi";
	process.env.PI_CODING_AGENT = "true";
	let runtime;
	let inventory = [];
	let stopThemeWatcher;
	let stage = "installed Pi loading";
	let published = false;
	const finish = async (value) => {
		if (published) return;
		published = true;
		if (options.sourcesOnly) {
			if (value.error) {
				console.error(value.error);
				return;
			}
			const result = value.sources;
			const n = (value) => value.toLocaleString("en-US");
			if (options.skillDescriptions) {
				const entries = result.descriptions;
				console.log(
					[
						`pi/agent: ${entries.length} visible skills, largest descriptions first`,
						"~tok  skill (raw description only, Pi chars/4)",
						...entries.map(
							(entry) => `${String(entry.tokens).padStart(4)}  ${entry.name}`,
						),
						"",
						"Largest five descriptions:",
						...entries
							.slice(0, 5)
							.map((entry) => `${entry.name}: ${entry.description}`),
					].join("\n"),
				);
				return;
			}
			console.log(
				[
					`APPEND_SYSTEM.md: ~${n(result.appendix)} tok`,
					`pi/agent skills (${result.piAgent.count}): ~${n(result.piAgent.tokens)} tok`,
					`Extension/package skills (${result.extensions.count}): ~${n(result.extensions.tokens)} tok`,
					`Shared catalog instructions/wrapper: ~${n(result.shared)} tok`,
					`Total: ~${n(result.total)} tok (Pi chars/4; group rounding may differ)`,
					`Project-local skills excluded: ${result.ignoredProjectSkills}; skill bodies and prompt templates excluded.`,
				].join("\n"),
			);
		} else {
			await publish(resultPath, value);
		}
	};
	try {
		const dist = await installedPiDist();
		const load = (path) => import(pathToFileURL(join(dist, path)).href);
		const [
			sdk,
			{ ReadOnlyAuthStorage },
			{ InMemoryCodingAgentModelsStore },
			{ builtInExtensions },
			trust,
			{ initTheme, stopThemeWatcher: stopWatcher },
			{ convertToLlm },
		] = await Promise.all([
			load("index.js"),
			load("core/auth-storage.js"),
			load("core/models-store.js"),
			load("extensions/index.js"),
			load("core/trust-manager.js"),
			load("modes/interactive/theme/theme.js"),
			load("core/messages.js"),
		]);
		stopThemeWatcher = stopWatcher;
		const { version } = JSON.parse(
			await readFile(join(dist, "../package.json"), "utf8"),
		);
		stage = "project trust resolution (use normal Pi to establish trust first)";
		let projectTrusted = false;
		if (!baseline) {
			projectTrusted =
				!trust.hasTrustRequiringProjectResources(cwd) ||
				new trust.ProjectTrustStore(agentDir).get(cwd);
			if (projectTrusted === null) {
				const policy = sdk.SettingsManager.create(cwd, agentDir, {
					projectTrusted: false,
				}).getDefaultProjectTrust();
				if (policy === "ask") throw new Error();
				projectTrusted = policy === "always";
			}
		}
		const settingsManager = baseline
			? sdk.SettingsManager.inMemory({}, { projectTrusted: false })
			: sdk.SettingsManager.create(cwd, agentDir, { projectTrusted });
		stage = "read-only authentication/model initialization";
		const modelRuntime = await sdk.ModelRuntime.create({
			credentials: new ReadOnlyAuthStorage(join(sourceAgentDir, "auth.json")),
			modelsPath: baseline ? null : join(agentDir, "models.json"),
			modelsStore: new InMemoryCodingAgentModelsStore(),
			allowModelNetwork: false,
		});
		let contextMode;
		let hasUI;
		let toolCalls = 0;
		let extensionErrors = 0;
		const observer = {
			name: "context-measurement",
			hidden: true,
			factory(pi) {
				pi.on("session_start", (_event, ctx) => {
					contextMode = ctx.mode;
					hasUI = ctx.hasUI;
				});
				// Keep schemas intact, but fail closed if the reply-only request tries tools.
				pi.on("tool_call", () => {
					toolCalls++;
					return {
						block: true,
						terminate: true,
						reason: "Context measurement permits no tool execution",
					};
				});
			},
		};
		stage = "resource/session initialization";
		runtime = await sdk.createAgentSessionRuntime(
			async ({ sessionManager, sessionStartEvent }) => {
				const services = await sdk.createAgentSessionServices({
					cwd,
					agentDir,
					settingsManager,
					modelRuntime,
					resourceLoaderOptions: {
						...resourceOptions(kind, [...builtInExtensions, observer], {
							cwd,
							agentDir,
							projectTrusted,
							extensionPaths: options.extensionPaths,
						}),
						...(options.inspectExtensions
							? {
									extensionsOverride(base) {
										inventory = inspectExtensionRegistrations(base);
										return base;
									},
								}
							: {}),
					},
				});
				if (
					services.diagnostics.some((d) => d.type === "error") ||
					services.resourceLoader.getExtensions().errors.length ||
					settingsManager.drainErrors().length ||
					modelRuntime.getError()
				)
					throw new Error();
				const model = selection
					? modelRuntime.getModel(selection.provider, selection.model)
					: undefined;
				if (selection && !model) {
					stage = baseline
						? "selected model unavailable in stock Pi (no user model config imported)"
						: "configured model selection";
					throw new Error();
				}
				const created = await sdk.createAgentSessionFromServices({
					services,
					sessionManager,
					sessionStartEvent,
					model,
					thinkingLevel: selection?.thinking,
				});
				return { ...created, services, diagnostics: services.diagnostics };
			},
			{ cwd, agentDir, sessionManager: sdk.SessionManager.inMemory(cwd) },
		);
		const { session } = runtime;
		stage = "startup model selection";
		if (!session.model) throw new Error();
		session.extensionRunner.onError(() => {
			extensionErrors++;
		});
		let requests = 0;
		let local;
		const messages = [];
		observeRequests(session.agent, (context) => {
			requests++;
			local = snapshot(context, sdk.estimateTokens);
		});
		session.subscribe((event) => {
			if (event.type === "message_end" && event.message.role === "assistant")
				messages.push(event.message);
		});
		const metadata = () => ({
			kind,
			cwd,
			version,
			provider: session.model?.provider,
			model: session.model?.id,
			thinking: session.thinkingLevel,
			contextWindow: session.model?.contextWindow,
			mode: contextMode,
			hasUI,
			projectTrusted,
			...(options.inspectExtensions ? { extensions: inventory } : {}),
		});
		const result = () => {
			if (requests !== 1 || toolCalls || extensionErrors) {
				return {
					kind,
					error: `Invalid run: ${requests} requests, ${toolCalls} tool calls, ${extensionErrors} extension errors`,
				};
			}
			try {
				const usage = summarize(
					messages
						.map((message) => JSON.stringify({ type: "message_end", message }))
						.join("\n"),
				);
				return { ...metadata(), local, ...usage, requests };
			} catch {
				return {
					kind,
					error:
						"No valid single response/usage; check provider access and read-only credential validity",
				};
			}
		};
		initTheme(settingsManager.getTheme(), false);
		stage = "headless startup hooks";
		await session.bindExtensions({
			mode: "json",
			onError: () => {
				extensionErrors++;
			},
		});
		if (contextMode !== "json" || hasUI || extensionErrors) throw new Error();
		if (options.inspectExtensions) {
			stage = "runtime extension inventory";
			inventory = inspectExtensionRegistrations(
				runtime.services.resourceLoader.getExtensions(),
				inventory,
			);
			if (
				options.extensionPaths !== undefined &&
				JSON.stringify(inventory.map((entry) => entry.path)) !==
					JSON.stringify(options.extensionPaths)
			) {
				stage = "explicit extension set/order validation";
				throw new Error();
			}
		}
		if (kind === "headless+0msg") {
			if (requests || messages.length) throw new Error();
			stage = "skill source attribution";
			const sources = options.sourcesOnly
				? await estimateSourceTokens(
						agentDir,
						sdk,
						runtime.services.resourceLoader.getSkills(),
					)
				: undefined;
			await finish({
				...metadata(),
				sources,
				local: snapshot(
					{
						...session.agent.state,
						messages: convertToLlm(session.messages),
					},
					sdk.estimateTokens,
				),
				requests,
			});
		} else {
			stage = "provider request (read-only credentials must already be valid)";
			await session.prompt(PROMPT);
			stage = "single-response/usage validation";
			await finish(result());
		}
	} catch {
		await finish({ kind, error: `${stage} failed` });
		process.exitCode = 1;
	} finally {
		for (const cleanup of [
			() => runtime?.dispose(),
			() => stopThemeWatcher?.(),
		]) {
			try {
				await cleanup();
			} catch {
				process.exitCode = 1;
			}
		}
	}
}

if (
	process.argv[1] &&
	import.meta.url === pathToFileURL(resolve(process.argv[1])).href
) {
	if (["--sources", "--skill-descriptions"].includes(process.argv[2])) {
		const agentDir =
			process.env.PI_CODING_AGENT_DIR ?? join(homedir(), ".pi", "agent");
		await runSession({
			kind: "headless+0msg",
			cwd: process.cwd(),
			agentDir,
			sourceAgentDir: agentDir,
			sourcesOnly: true,
			skillDescriptions: process.argv[2] === "--skill-descriptions",
		});
	} else {
		await runSession(JSON.parse(await readFile(process.argv[2], "utf8")));
	}
	process.exit(process.exitCode ?? 0);
}