Luigit
repositories / pi-ext

pi-ext

bugabingas pi extensions

owned by admin

extensions/klaus/index.ts

Raw
import { lazyStream } from "@earendil-works/pi-ai";
import type {
	ExtensionAPI,
	ExtensionContext,
	ProviderModelConfig,
} from "@earendil-works/pi-coding-agent";
import type { QueryCoordinator } from "./src/coordinator.js";
import { closeDebug, dbg, debugSpan } from "./src/debug.js";
import {
	modelsSignature,
	projectBuiltinModels,
	projectModels,
} from "./src/models.js";
import { KLAUS_API } from "./src/protocol.js";

dbg?.("index.module");

/** Pi's catalog carries user overrides, so the session registry outranks the
 * builtin catalog whenever it still knows Claude models. */
function sessionModels(activeContext: ExtensionContext): ProviderModelConfig[] {
	const models = projectModels(activeContext.modelRegistry.getAvailable());
	return models.length ? models : projectBuiltinModels();
}

export default function klaus(pi: ExtensionAPI): void | Promise<void> {
	dbg?.("index.extension.start");
	let context: ExtensionContext | undefined;
	let warnedFable = false;
	let warnedPromptDrift = false;
	let coordinatorPromise: Promise<QueryCoordinator> | undefined;
	const getCoordinator = () =>
		(coordinatorPromise ??= import("./src/coordinator.js").then(
			({ QueryCoordinator }) =>
				new QueryCoordinator(
					async (modelId) => {
						const finish = debugSpan?.("index.oauth");
						try {
							if (!context)
								throw new Error("Klaus has no active Pi session context.");
							const source = context.modelRegistry.find("anthropic", modelId);
							dbg?.("index.oauth.model", { found: Boolean(source) });
							if (!source)
								throw new Error(
									`Anthropic model ${modelId} is unavailable in Pi.`,
								);
							const usingOAuth = context.modelRegistry.isUsingOAuth(source);
							dbg?.("index.oauth.kind", { usingOAuth });
							if (!usingOAuth) {
								throw new Error(
									"Klaus requires Anthropic OAuth. Use Pi /login for Anthropic first.",
								);
							}
							const auth =
								await context.modelRegistry.getProviderAuth("anthropic");
							const token = auth?.auth.apiKey;
							dbg?.("index.oauth.resolved", { hasAuth: Boolean(auth) });
							if (!token)
								throw new Error(
									"Pi could not resolve Anthropic OAuth for Klaus.",
								);
							if (modelId === "claude-fable-5-1" && !warnedFable) {
								dbg?.("index.oauth.fableWarning");
								warnedFable = true;
								context.ui.notify(
									"Klaus Fable may consume Anthropic usage credits when subscription discovery omits it.",
									"warning",
								);
							}
							finish?.();
							return token;
						} catch (error) {
							finish?.("error", { kind: "oauth" });
							throw error;
						}
					},
					() => {
						const cwd = context?.cwd ?? process.cwd();
						dbg?.("index.currentCwd");
						return cwd;
					},
					() => {
						const scope = context
							? {
									id: context.sessionManager.getSessionId(),
									persisted:
										context.sessionManager.getSessionFile() !== undefined,
									dir: context.sessionManager.getSessionDir(),
								}
							: undefined;
						dbg?.("index.currentSession", {
							hasScope: Boolean(scope),
							persisted: scope?.persisted,
						});
						return scope;
					},
					(message) => {
						dbg?.("index.promptDriftWarning", {
							alreadyWarned: warnedPromptDrift,
						});
						if (warnedPromptDrift || !context) return;
						warnedPromptDrift = true;
						context.ui.notify(message, "warning");
					},
				),
		));

	pi.registerCommand("klaus-setup-token", {
		description: "Mint a long-lived Claude OAuth token with your subscription",
		async handler(_args, commandCtx) {
			const { runSetupTokenCommand } = await import("./src/setup-token.js");
			await runSetupTokenCommand(commandCtx);
		},
	});

	let preparedSession: string | undefined;
	let preparation: Promise<void> | undefined;
	const readyCoordinator = async () => {
		const coordinator = await getCoordinator();
		const active = context;
		const session = active
			? `${active.sessionManager.getSessionId()}\0${active.sessionManager.getSessionDir()}`
			: "none";
		if (preparedSession === session) return coordinator;
		preparation ??= coordinator
			.preparePrimaryCache(
				active
					? {
							id: active.sessionManager.getSessionId(),
							persisted: active.sessionManager.getSessionFile() !== undefined,
							dir: active.sessionManager.getSessionDir(),
						}
					: { id: "", persisted: false, dir: process.cwd() },
			)
			.then(() => {
				preparedSession = session;
			})
			.finally(() => {
				preparation = undefined;
			});
		await preparation;
		return coordinator;
	};

	let registeredSignature: string | undefined;
	const register = (models: ProviderModelConfig[]): void => {
		const signature = modelsSignature(models);
		if (signature === registeredSignature) {
			dbg?.("index.register.unchanged", { count: models.length });
			return;
		}
		registeredSignature = signature;
		dbg?.("index.register");
		pi.registerProvider("klaus", {
			name: "Klaus",
			baseUrl: "https://klaus.invalid",
			apiKey: "klaus",
			authHeader: false,
			api: KLAUS_API,
			models,
			streamSimple: (model, requestContext, options) => {
				dbg?.("index.streamSimple", {
					messageCount: requestContext.messages.length,
				});
				return lazyStream(model, async () =>
					(await readyCoordinator()).stream(model, requestContext, options),
				);
			},
		});
	};

	register(projectBuiltinModels());

	pi.on("session_start", (_event, nextContext) => {
		dbg?.("index.session_start", {
			persisted: nextContext.sessionManager.getSessionFile() !== undefined,
		});
		context = nextContext;
		preparedSession = undefined;
		warnedFable = false;
		warnedPromptDrift = false;
		register(sessionModels(nextContext));
	});
	pi.on("turn_start", (_event, activeContext) => {
		// Pi has no catalog-reload event, so a turn boundary is the cheapest
		// place to notice a refreshed models.json. Registration is skipped when
		// the projection is unchanged.
		register(sessionModels(activeContext));
	});
	pi.on("turn_end", async (_event, activeContext) => {
		const leafId = activeContext.sessionManager.getLeafId();
		dbg?.("index.turn_end");
		const coordinator = await coordinatorPromise?.catch(() => undefined);
		if (coordinator && leafId)
			await coordinator.persistCheckpoint(
				activeContext.sessionManager.getSessionId(),
				leafId,
			);
	});
	const flushAndClose = async (
		activeContext: ExtensionContext,
		reason: string,
	): Promise<void> => {
		const coordinator = await coordinatorPromise?.catch(() => undefined);
		if (!coordinator) return;
		const leafId = activeContext.sessionManager.getLeafId();
		const finish = debugSpan?.("index.flush", { hasLeaf: Boolean(leafId) });
		try {
			if (leafId)
				await coordinator.persistCheckpoint(
					activeContext.sessionManager.getSessionId(),
					leafId,
				);
			await coordinator.closeAll(reason);
			finish?.();
		} catch (error) {
			finish?.("error", { kind: "flush" });
			throw error;
		}
	};
	pi.on("session_before_switch", async (_event, activeContext) => {
		dbg?.("index.session_before_switch");
		await flushAndClose(activeContext, "Pi switched sessions.");
	});
	pi.on("session_before_fork", async (_event, activeContext) => {
		dbg?.("index.session_before_fork");
		await flushAndClose(activeContext, "Pi forked the session.");
	});
	pi.on("session_before_compact", async (_event, activeContext) => {
		dbg?.("index.session_before_compact");
		await flushAndClose(activeContext, "Pi compacted the session.");
	});
	pi.on("session_shutdown", async (_event, activeContext) => {
		dbg?.("index.session_shutdown");
		try {
			await flushAndClose(activeContext, "Pi shut down the session.");
		} finally {
			context = undefined;
			preparedSession = undefined;
			closeDebug();
		}
	});
}