Luigit
repositories / pi-ext

pi-ext

bugabingas pi extensions

owned by admin

extensions/ultra/drivers.ts

Raw
// ultra — real SDK drivers (Task 9).
//
// The runner's SDK capabilities, wrapped behind injected factories so option
// construction remains unit-testable offline. `implementation.ts` supplies the
// real `createAgentSession`, `SessionManager`, and `DefaultResourceLoader`.
//
// SDK shapes verified against the installed declarations:
//   - createAgentSession resolves to `{ session }`.
//   - DefaultResourceLoader supports isolated resource discovery.
//   - modelRegistry.find(provider, id) resolves explicit model refs.

import { existsSync } from "node:fs";
import { resolve } from "node:path";
import {
	BUILTIN_TOOL_NAMES,
	ULTRA_RESERVED_TOOL_NAMES,
} from "./dynamic-contract.ts";
import type { ResolvedDynamicExtension } from "./dynamic-types.ts";
import type { SubagentSessionRef } from "./journal.ts";
import type {
	ResolveModel,
	SessionFactory,
	SessionLike,
	SubagentSessionHandle,
} from "./runner.ts";

export { BUILTIN_TOOL_NAMES } from "./dynamic-contract.ts";

// ---------------------------------------------------------------------------
// sessionFactory — wraps createAgentSession with an isolated resource loader
// ---------------------------------------------------------------------------

/** The subset of the SDK the session factory needs (injected for testing). */
export interface HeadlessSession extends SessionLike {
	setActiveToolsByName(toolNames: string[]): void;
	getActiveToolNames(): string[];
	bindExtensions(bindings: { mode: "print" }): Promise<void>;
}

export type FlagValue = boolean | string;

export interface SessionFactoryDeps {
	createAgentSession: (
		opts: Record<string, unknown>,
	) => Promise<{ session: HeadlessSession }> | { session: HeadlessSession };
	DefaultResourceLoader: new (
		opts: Record<string, unknown>,
	) => {
		reload(): Promise<void>;
		getExtensions?(): unknown;
	};
	getAgentDir: () => string;
	cwd: string;
	/** When true, sub-agents load every host extension. */
	loadExtensions?: boolean;
	/** Explicit extensions loaded when broad discovery is disabled. */
	extensionPaths?: string[];
	/** When true, sub-agents discover host skills. */
	loadSkills?: boolean;
	/**
	 * Parent CLI extension flags, as Pi's `parseArgs(...).unknownFlags` returns them.
	 * Applied like Pi's own CLI: only flags a child extension registers.
	 */
	inheritedFlags?: ReadonlyMap<string, FlagValue>;
}

export function makeSessionFactory(deps: SessionFactoryDeps): SessionFactory {
	return async ({
		model,
		thinkingLevel,
		tools,
		customTools,
		extensionPaths,
		dynamicExtensions,
		appendSystemPrompt,
		sessionManager,
		flags,
	}) => {
		// resourceLoader is REQUIRED for controlled discovery (P1). reload() MUST
		// complete before the session is constructed.
		const additionalExtensionPaths = [
			...(deps.extensionPaths ?? []),
			...(extensionPaths ?? []),
		];
		const resourceLoader = new deps.DefaultResourceLoader({
			cwd: deps.cwd,
			agentDir: deps.getAgentDir(),
			noExtensions: !deps.loadExtensions,
			additionalExtensionPaths:
				additionalExtensionPaths.length > 0
					? additionalExtensionPaths
					: undefined,
			noSkills: !deps.loadSkills,
			noPromptTemplates: true,
			noThemes: true,
			// Applicable global/project instructions belong in every child prompt.
			// Conversation isolation is provided by the child session manager.
			noContextFiles: false,
			...(appendSystemPrompt?.length
				? {
						appendSystemPromptOverride: (base: string[]) => [
							...base,
							...appendSystemPrompt,
						],
					}
				: {}),
		});
		await resourceLoader.reload();
		const loaded = resourceLoader.getExtensions?.() as
			| LoadedExtensionsResult
			| undefined;
		if (dynamicExtensions?.length && loaded)
			applyDeclaredOverrides(
				loaded,
				dynamicExtensions,
				new Set(customTools.map((tool) => tool.name)),
			);
		// Flags must be set before bindExtensions emits session_start.
		applyFlags(loaded, deps.inheritedFlags, flags);

		const activeTools =
			deps.loadSkills && !tools.includes("read") ? [...tools, "read"] : tools;
		const { session } = await deps.createAgentSession({
			cwd: deps.cwd,
			model,
			resourceLoader,
			thinkingLevel,
			sessionManager, // clean child, or a matching interrupted child's manager
			customTools, // [structuredOutputTool] — built by the runner from the step schema
		});
		// Passing `tools` to createAgentSession creates a hard registry allowlist,
		// which prevents loaded extensions from activating their own tools.
		session.setActiveToolsByName(activeTools);
		await session.bindExtensions({ mode: "print" });
		// Extensions may switch requested tools off at session_start; say which flag owns them.
		const active = new Set(session.getActiveToolNames());
		const inactive = activeTools.filter((name) => !active.has(name));
		if (inactive.length > 0) {
			session.dispose();
			throw new Error(inactiveToolsMessage(inactive, loaded));
		}
		return session;
	};
}

interface RegisteredFlag {
	name: string;
	type: "boolean" | "string";
	description?: string;
}

function registeredFlags(
	loaded: LoadedExtensionsResult | undefined,
): Map<string, RegisteredFlag> {
	const flags = new Map<string, RegisteredFlag>();
	for (const extension of loaded?.extensions ?? [])
		for (const [name, flag] of extension.flags ?? [])
			if (!flags.has(name)) flags.set(name, flag as RegisteredFlag);
	return flags;
}

/** Mirror Pi's CLI: inherited flags apply only when registered; step flags must be. */
function applyFlags(
	loaded: LoadedExtensionsResult | undefined,
	inherited: ReadonlyMap<string, FlagValue> | undefined,
	step: Record<string, FlagValue> | undefined,
): void {
	if (!inherited?.size && !step) return;
	const values = loaded?.runtime?.flagValues;
	const registered = registeredFlags(loaded);
	for (const [name, value] of inherited ?? []) {
		const flag = registered.get(name);
		if (!flag || !values) continue;
		// Pi's parser cannot know flag types, so a boolean flag may have captured a value.
		if (flag.type === "boolean") values.set(name, true);
		else if (typeof value === "string") values.set(name, value);
	}
	for (const [name, value] of Object.entries(step ?? {})) {
		const flag = registered.get(name);
		if (!flag || !values)
			throw new Error(
				`ultra: step flag --${name} is not registered by any sub-agent extension.`,
			);
		if (typeof value !== flag.type)
			throw new Error(
				`ultra: step flag --${name} expects a ${flag.type} value.`,
			);
		values.set(name, value);
	}
}

export interface SubagentFlagEntry {
	extension: string;
	tools: string[];
	flags: RegisteredFlag[];
}

/**
 * Load the sub-agent extension set once, without a session, and list the flags
 * of every extension that also owns tools. Uses only loaded-extension metadata.
 */
export async function loadSubagentFlagCatalog(
	deps: Pick<
		SessionFactoryDeps,
		| "DefaultResourceLoader"
		| "getAgentDir"
		| "cwd"
		| "loadExtensions"
		| "extensionPaths"
	>,
): Promise<SubagentFlagEntry[]> {
	const loader = new deps.DefaultResourceLoader({
		cwd: deps.cwd,
		agentDir: deps.getAgentDir(),
		noExtensions: !deps.loadExtensions,
		additionalExtensionPaths: deps.extensionPaths?.length
			? deps.extensionPaths
			: undefined,
		noSkills: true,
		noPromptTemplates: true,
		noThemes: true,
		noContextFiles: true,
	});
	await loader.reload();
	const loaded = loader.getExtensions?.() as LoadedExtensionsResult | undefined;
	return (loaded?.extensions ?? [])
		.filter((extension) => extension.tools?.size && extension.flags?.size)
		.map((extension) => ({
			extension: extensionLabel(extension.path),
			tools: [...(extension.tools?.keys() ?? [])],
			flags: [...(extension.flags?.values() ?? [])] as RegisteredFlag[],
		}));
}

/** Orchestrator guidance: which step.flags switch on which extension's tools. */
export function subagentFlagPrompt(catalog: SubagentFlagEntry[]): string {
	return [
		"# ultra sub-agent flags",
		"Some extensions keep their tools off until a flag is set; a step listing such a tool fails unless step.flags enables it, for example flags: { nushell: true }.",
		"Children inherit the parent's CLI flags; step.flags override them. Rows: extension: tools; flags (type): description.",
		...catalog.map(
			(entry) =>
				`- ${entry.extension}: ${entry.tools.join(", ")}; ${entry.flags
					.map(
						(flag) =>
							`--${flag.name} (${flag.type})${flag.description ? `: ${flag.description}` : ""}`,
					)
					.join("; ")}`,
		),
	].join("\n");
}

/** Directory name for `<dir>/index.ts` entrypoints, else the file name. */
export function extensionLabel(path: string): string {
	const parts = path.replaceAll("\\", "/").split("/");
	const file = parts.at(-1) ?? path;
	return /^index\.[cm]?[jt]s$/u.test(file) ? (parts.at(-2) ?? file) : file;
}

function inactiveToolsMessage(
	inactive: string[],
	loaded: LoadedExtensionsResult | undefined,
): string {
	const details = inactive.map((name) => {
		const owner = loaded?.extensions.find((extension) =>
			extension.tools?.has(name),
		);
		if (!owner) return name;
		const flags = [...(owner.flags?.values() ?? [])] as RegisteredFlag[];
		const hint = flags.length
			? `; flags: ${flags.map((flag) => `--${flag.name}${flag.description ? ` (${flag.description})` : ""}`).join(", ")}`
			: "";
		return `${name} (extension ${extensionLabel(owner.path)}${hint})`;
	});
	return `ultra: requested tools were switched off by their extensions at startup: ${details.join("; ")}. Enable them through the owning extension's flag in step.flags.`;
}

interface LoadedExtension {
	path: string;
	resolvedPath?: string;
	tools?: Map<string, unknown>;
	commands?: Map<string, unknown>;
	flags?: Map<string, unknown>;
	shortcuts?: Map<string, unknown>;
	messageRenderers?: Map<string, unknown>;
	entryRenderers?: Map<string, unknown>;
}

interface LoadedExtensionsResult {
	extensions: LoadedExtension[];
	runtime?: {
		flagValues?: Map<string, FlagValue>;
		pendingProviderRegistrations?: Array<{
			name: string;
			extensionPath: string;
		}>;
		pendingNativeProviderRegistrations?: Array<{
			provider: { name?: string; id?: string };
			extensionPath: string;
		}>;
	};
}

function normalizedPath(path: string): string {
	const absolute = resolve(path).replaceAll("\\", "/");
	return process.platform === "win32" ? absolute.toLowerCase() : absolute;
}

function applyDeclaredOverrides(
	loaded: unknown,
	dynamicExtensions: ResolvedDynamicExtension[],
	sdkToolNames: Set<string>,
): void {
	const result = loaded as LoadedExtensionsResult;
	if (!Array.isArray(result?.extensions)) return;
	const dynamicByPath = new Map(
		dynamicExtensions.map((extension) => [
			normalizedPath(extension.loadPath),
			extension,
		]),
	);
	const registrations: Array<[keyof LoadedExtension, string]> = [
		["tools", "tool"],
		["commands", "command"],
		["flags", "flag"],
		["shortcuts", "shortcut"],
		["messageRenderers", "renderer"],
		["entryRenderers", "entry-renderer"],
	];
	const seen = new Map<string, Map<string, unknown>>();
	for (const loadedExtension of result.extensions) {
		const extension = dynamicByPath.get(
			normalizedPath(loadedExtension.resolvedPath ?? loadedExtension.path),
		);
		for (const [field, qualifier] of registrations) {
			const values = loadedExtension[field];
			if (!(values instanceof Map)) continue;
			for (const name of [...values.keys()]) {
				const qualified = `${qualifier}:${name}`;
				if (
					extension &&
					qualifier === "tool" &&
					(ULTRA_RESERVED_TOOL_NAMES.has(name) || sdkToolNames.has(name))
				)
					throw new Error(
						`ultra: dynamic extension "${extension.name}" cannot register reserved ${qualified}.`,
					);
				if (
					extension &&
					qualifier === "tool" &&
					BUILTIN_TOOL_NAMES.has(name) &&
					!extension.overrides.includes(qualified)
				)
					throw new Error(
						`ultra: dynamic extension "${extension.name}" collides with ${qualified}; declare that exact override in manifest.json or rename it.`,
					);
				const prior = seen.get(qualified);
				if (prior && extension) {
					if (!extension.overrides.includes(qualified))
						throw new Error(
							`ultra: dynamic extension "${extension.name}" collides with ${qualified}; declare that exact override in manifest.json or rename it.`,
						);
					prior.delete(name);
				}
				if (!prior) seen.set(qualified, values);
			}
		}
	}
	applyProviderOverrides(result, dynamicExtensions);
}

function applyProviderOverrides(
	result: LoadedExtensionsResult,
	dynamicExtensions: ResolvedDynamicExtension[],
): void {
	const dynamicByPath = new Map(
		dynamicExtensions.map((extension) => [
			normalizedPath(extension.loadPath),
			extension,
		]),
	);
	for (const registrations of [
		result.runtime?.pendingProviderRegistrations,
		result.runtime?.pendingNativeProviderRegistrations,
	]) {
		if (!registrations) continue;
		const seen = new Map<string, number>();
		for (let index = 0; index < registrations.length; index++) {
			const registration = registrations[index];
			const name =
				"name" in registration
					? registration.name
					: (registration.provider.id ?? registration.provider.name);
			if (!name) continue;
			const prior = seen.get(name);
			const extension = dynamicByPath.get(
				normalizedPath(registration.extensionPath),
			);
			if (prior !== undefined && extension) {
				const qualified = `provider:${name}`;
				if (!extension.overrides.includes(qualified))
					throw new Error(
						`ultra: dynamic extension "${extension.name}" collides with ${qualified}; declare that exact override in manifest.json or rename it.`,
					);
				registrations.splice(prior, 1);
				index--;
				seen.clear();
				for (let i = 0; i <= index; i++) {
					const priorRegistration = registrations[i];
					const priorName =
						"name" in priorRegistration
							? priorRegistration.name
							: (priorRegistration.provider.id ??
								priorRegistration.provider.name);
					if (priorName) seen.set(priorName, i);
				}
				continue;
			}
			seen.set(name, index);
		}
	}
}

// ---------------------------------------------------------------------------
// subagentSessionManager — clean persistent child or matching continuation
// ---------------------------------------------------------------------------

interface SessionManagerLike {
	getSessionId(): string;
	getSessionFile(): string | undefined;
	appendSessionInfo(name: string): string;
}

export interface SubagentSessionManagerDeps {
	SessionManager: {
		create(
			cwd: string,
			sessionDir?: string,
			options?: { parentSession?: string },
		): SessionManagerLike;
		open(
			path: string,
			sessionDir?: string,
			cwdOverride?: string,
		): SessionManagerLike;
		inMemory(cwd: string): SessionManagerLike;
	};
	cwd: string;
	parentSession?: string;
	retention: "all" | "none";
	/** Parent session's actual storage directory. */
	sessionDir?: string;
}

export interface SubagentSessionManagerArgs {
	name: string;
	retained?: SubagentSessionRef;
}

export interface ManagedSubagentSession extends SubagentSessionHandle {
	ref?: SubagentSessionRef;
}

/** Persist only beneath a persisted parent. A missing retained file starts clean. */
export function makeSubagentSessionManager(
	deps: SubagentSessionManagerDeps,
): (args: SubagentSessionManagerArgs) => ManagedSubagentSession {
	return ({ name, retained }) => {
		if (deps.retention === "none" || !deps.parentSession) {
			const manager = deps.SessionManager.inMemory(deps.cwd);
			return {
				manager,
				id: manager.getSessionId(),
				resumed: false,
			};
		}

		if (retained && existsSync(retained.path)) {
			try {
				const manager = deps.SessionManager.open(
					retained.path,
					deps.sessionDir,
					deps.cwd,
				);
				if (manager.getSessionId() === retained.id)
					return {
						manager,
						id: manager.getSessionId(),
						resumed: true,
						ref: retained,
					};
			} catch {}
		}

		const manager = deps.SessionManager.create(deps.cwd, deps.sessionDir, {
			parentSession: deps.parentSession,
		});
		manager.appendSessionInfo(name);
		const path = manager.getSessionFile();
		if (!path) return { manager, id: manager.getSessionId(), resumed: false };
		return {
			manager,
			id: manager.getSessionId(),
			resumed: false,
			ref: { id: manager.getSessionId(), path },
		};
	};
}

// ---------------------------------------------------------------------------
// resolveModel — "provider/model" ref → Model object, else the host default
// ---------------------------------------------------------------------------

/** Split a `provider/model` ref on the first slash. */
function parseModelRef(ref: string): { provider: string; model: string } {
	const i = ref.indexOf("/");
	if (i <= 0 || i === ref.length - 1) {
		throw new Error(
			`ultra: invalid model ref "${ref}" — use "provider/model".`,
		);
	}
	return { provider: ref.slice(0, i), model: ref.slice(i + 1) };
}

export interface ResolveModelDeps {
	modelRegistry: { find(provider: string, modelId: string): unknown };
	/** The host's current model (`ctx.model`) — the default when a step omits `model`. */
	defaultModel: unknown;
}

export function makeResolveModel(deps: ResolveModelDeps): ResolveModel {
	return (modelRef) => {
		if (!modelRef) return deps.defaultModel;
		const { provider, model } = parseModelRef(modelRef);
		const found = deps.modelRegistry.find(provider, model);
		if (!found) throw new Error(`ultra: model not found: ${modelRef}`);
		return found;
	};
}