Luigit
repositories / pi-ext

pi-ext

bugabingas pi extensions

owned by admin

extensions/klaus/src/agent-sdk.ts

Raw
import {
	type EffortLevel,
	type Query,
	query,
	type SDKMessage,
	type SDKUserMessage,
	type ThinkingConfig,
} from "@anthropic-ai/claude-agent-sdk";
import type { JsonObject } from "@earendil-works/pi-ai";
import { dbg, diagnosticKind } from "./debug.js";
import type {
	KlausQueryCallbacks,
	KlausRequest,
	KlausTool,
} from "./protocol.js";
import { CLAUDE_VERSION, SDK_VERSION } from "./protocol.js";
import { createChildRuntime } from "./runtime/index.js";
import { TOOL_ID_FIELD, ToolBridge } from "./tool-bridge.js";

interface BlockState {
	type: "text" | "thinking" | "tool";
	id?: string;
	name?: string;
	text: string;
	signature?: string;
	redacted?: boolean;
}

interface MessageState {
	latestPosition?: string;
	responseId?: string;
	responseModel?: string;
	sdkSessionId?: string;
	deferredOutputLimit?: Parameters<KlausQueryCallbacks["onResult"]>[0];
	usage: {
		input: number;
		output: number;
		cacheRead: number;
		cacheWrite: number;
	};
	reportedUsage: {
		input: number;
		output: number;
		cacheRead: number;
		cacheWrite: number;
	};
}

export interface KlausTurn {
	prompt: string;
	images?: KlausRequest["images"];
}

export interface KlausQueryHandle {
	bridge: ToolBridge;
	close(reason: string): Promise<void>;
	/** True while the child is alive, has no turn in flight, and can accept the
	 * next Pi turn through its open input stream. */
	isIdle(): boolean;
	/** Send the next Pi turn into the running child instead of paying another
	 * spawn and initialize round trip. */
	send(turn: KlausTurn, callbacks: KlausQueryCallbacks): void;
	/** Replace the tool list of an idle child so a Pi tool-list change does not
	 * cost a canonical replay. */
	setTools(tools: KlausRequest["tools"]): Promise<void>;
}

interface PromptChannel {
	stream: AsyncGenerator<SDKUserMessage>;
	push(message: SDKUserMessage): void;
	end(): void;
}

/** Streaming input mode keeps the child waiting for the next turn, so Klaus
 * owns the message stream instead of yielding one prompt and ending it. */
function createPromptChannel(first: SDKUserMessage): PromptChannel {
	const queue: SDKUserMessage[] = [first];
	let notify: (() => void) | undefined;
	let ended = false;
	async function* stream(): AsyncGenerator<SDKUserMessage> {
		while (true) {
			const next = queue.shift();
			if (next) {
				dbg?.("sdk.prompt.yield");
				yield next;
				continue;
			}
			if (ended) return;
			await new Promise<void>((resolve) => {
				notify = resolve;
			});
		}
	}
	const wake = (): void => {
		const pending = notify;
		notify = undefined;
		pending?.();
	};
	return {
		stream: stream(),
		push(message) {
			queue.push(message);
			wake();
		},
		end() {
			ended = true;
			wake();
		},
	};
}

function effort(level: KlausRequest["thinking"]): EffortLevel | undefined {
	if (!level) return undefined;
	return level === "minimal" ? "low" : level;
}

/** Fallback budgets mirror Pi's DEFAULT_THINKING_BUDGETS so a Haiku request
 * without Pi-configured budgets behaves like native Pi usage. */
function defaultThinkingBudget(level: KlausRequest["thinking"]): number {
	switch (level) {
		case "minimal":
			return 1024;
		case "low":
			return 2048;
		case "medium":
			return 8192;
		default:
			return 16384;
	}
}

/** Adaptive-thinking models always run adaptive with effort, matching Pi's
 * native treatment of them: Pi thinking budgets are not a lever for these
 * models and must not pin them into fixed-budget thinking. Only budget-based
 * models (Haiku) consume an explicit budget. */
export function klausThinkingConfig(
	request: Pick<KlausRequest, "modelId" | "thinking" | "thinkingBudget">,
): ThinkingConfig {
	if (!request.thinking) return { type: "disabled" };
	if (request.modelId === "claude-haiku-4-5") {
		return {
			type: "enabled",
			budgetTokens:
				request.thinkingBudget ?? defaultThinkingBudget(request.thinking),
		};
	}
	return { type: "adaptive" };
}

function errorText(error: unknown): string {
	return error instanceof Error ? error.message : String(error);
}

function redact(text: string, secret: string): string {
	return secret ? text.replaceAll(secret, "<redacted>") : text;
}

/** The API rejects an invalid tool schema by position, and Claude Code cannot
 * pre-filter those tools because Klaus disables feature-flag fetching. Klaus
 * knows the order it sent, so it names the tool instead. */
function nameToolSchemaError(text: string, tools: KlausTool[]): string {
	const match = /tools\.(\d+)\.custom\.input_schema/.exec(text);
	if (!match) return text;
	const position = Number(match[1]);
	const tool = Number.isSafeInteger(position) ? tools[position] : undefined;
	dbg?.("sdk.toolSchemaError");
	return tool
		? `Klaus tool "${tool.name}" has an input schema the Anthropic API rejects: ${text}`
		: `A Klaus tool at position ${position} has an input schema the Anthropic API rejects: ${text}`;
}

function normalizeProviderError(text: string, tools: KlausTool[]): string {
	const named = nameToolSchemaError(text, tools);
	return /context(?:_| )length|prompt is too long|too many tokens/i.test(named)
		? `context_length_exceeded: ${named}`
		: named;
}

function isOutputLimit(text: string): boolean {
	return /max_output_tokens|response exceeded the \d+ output token maximum/i.test(
		text,
	);
}

interface ProviderUsage {
	input_tokens: number;
	output_tokens: number;
	cache_read_input_tokens?: number;
	cache_creation_input_tokens?: number;
}

function remainingUsage(
	usage: ProviderUsage,
	reported: MessageState["reportedUsage"],
): MessageState["usage"] {
	return {
		input: Math.max(0, usage.input_tokens - reported.input),
		output: Math.max(0, usage.output_tokens - reported.output),
		cacheRead: Math.max(
			0,
			(usage.cache_read_input_tokens ?? 0) - reported.cacheRead,
		),
		cacheWrite: Math.max(
			0,
			(usage.cache_creation_input_tokens ?? 0) - reported.cacheWrite,
		),
	};
}

function addUsage(
	target: MessageState["reportedUsage"],
	usage: MessageState["usage"],
): void {
	target.input += usage.input;
	target.output += usage.output;
	target.cacheRead += usage.cacheRead;
	target.cacheWrite += usage.cacheWrite;
}

function originalToolName(name: string): string {
	return name.startsWith("mcp__klaus__")
		? name.slice("mcp__klaus__".length)
		: name;
}

type ImageMediaType = "image/jpeg" | "image/png" | "image/gif" | "image/webp";

function imageMediaType(value: string): ImageMediaType {
	if (
		value === "image/jpeg" ||
		value === "image/png" ||
		value === "image/gif" ||
		value === "image/webp"
	) {
		return value;
	}
	throw new Error(`Unsupported Klaus image type ${value}.`);
}

function userMessage(turn: KlausTurn): SDKUserMessage {
	dbg?.("sdk.prompt.build", { imageCount: turn.images?.length ?? 0 });
	return {
		type: "user",
		message: {
			role: "user",
			content: [
				{ type: "text", text: turn.prompt },
				...(turn.images ?? []).map((image) => ({
					type: "image" as const,
					source: {
						type: "base64" as const,
						media_type: imageMediaType(image.mimeType),
						data: image.data,
					},
				})),
			],
		},
		parent_tool_use_id: null,
		origin: { kind: "human" },
	};
}

export async function startSdkQuery(
	request: KlausRequest,
	oauthToken: string,
	callbacks: KlausQueryCallbacks,
): Promise<KlausQueryHandle> {
	dbg?.("sdk.start", { imageCount: request.images?.length ?? 0 });
	for (const image of request.images ?? []) imageMediaType(image.mimeType);
	const runtime = await createChildRuntime(
		oauthToken,
		request.headers,
		request.env,
		request.maxTokens,
	);
	dbg?.("sdk.runtime.ready");
	let active = callbacks;
	let bridge = new ToolBridge(request.tools, () => active.onActivity());
	let bridgeGeneration = 0;
	dbg?.("sdk.bridge.ready");
	const abortController = new AbortController();
	const blocks = new Map<number, BlockState>();
	const channel = createPromptChannel(userMessage(request));
	let queryHandle: Query | undefined;
	let pumpPromise: Promise<void> | undefined;
	let cleanupPromise: Promise<void> | undefined;
	let closePromise: Promise<void> | undefined;
	let closed = false;
	let turnActive = true;
	let reusable = false;
	let stderr = "";
	const messageState: MessageState = {
		usage: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 },
		reportedUsage: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 },
	};
	/** Every turn reports its own usage and content indices, so no per-turn
	 * accounting may leak into the next turn on the same child. */
	const startTurn = (): void => {
		blocks.clear();
		messageState.usage = { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 };
		messageState.reportedUsage = {
			input: 0,
			output: 0,
			cacheRead: 0,
			cacheWrite: 0,
		};
		messageState.deferredOutputLimit = undefined;
		messageState.responseId = undefined;
		messageState.responseModel = undefined;
		turnActive = true;
		reusable = false;
	};

	/** Claude only serves control requests such as `interrupt()` in streaming
	 * input mode, so every query uses the generator prompt. */
	const interruptDeadlineMs = 1_500;
	const gracefulStop = async (): Promise<void> => {
		if (!queryHandle || !turnActive) return;
		dbg?.("sdk.interrupt.start");
		await Promise.race([
			queryHandle.interrupt().catch(() => undefined),
			new Promise((resolve) => setTimeout(resolve, interruptDeadlineMs)),
		]);
		dbg?.("sdk.interrupt.end");
	};

	const cleanup = (reason: string): Promise<void> => {
		dbg?.("sdk.cleanup.request", { existing: Boolean(cleanupPromise) });
		cleanupPromise ??= (async () => {
			dbg?.("sdk.cleanup.start");
			await bridge.close(reason).catch(() => undefined);
			await runtime.cleanup();
			dbg?.("sdk.cleanup.end");
		})();
		return cleanupPromise;
	};
	const close = (reason: string): Promise<void> => {
		dbg?.("sdk.close.request", { existing: Boolean(closePromise), closed });
		closePromise ??= (async () => {
			if (!closed) {
				dbg?.("sdk.close.abort");
				await gracefulStop();
				closed = true;
				reusable = false;
				channel.end();
				abortController.abort(reason);
				queryHandle?.close();
			}
			await pumpPromise?.catch(() => undefined);
			await cleanup(reason);
		})();
		return closePromise;
	};

	try {
		dbg?.("sdk.query.call");
		queryHandle = query({
			prompt: channel.stream,
			options: {
				abortController,
				cwd: request.cwd,
				env: runtime.env,
				pathToClaudeCodeExecutable: runtime.executable,
				model: request.selector,
				systemPrompt: request.systemPrompt,
				tools: [],
				skills: [],
				settingSources: [],
				strictMcpConfig: true,
				includePartialMessages: true,
				persistSession: true,
				sessionStore: request.sessionStore,
				/** Klaus deletes the child config directory when the query ends, so
				 * the mirrored transcript is the only durable copy and must not sit
				 * in a batch when a turn is interrupted. */
				sessionStoreFlush: "eager",
				resume: request.resume,
				resumeSessionAt: request.resumeAt,
				forkSession: request.forkSession,
				permissionMode: "default",
				thinking: klausThinkingConfig(request),
				effort: effort(request.thinking),
				betas: runtime.betas,
				settings: {
					autoCompactEnabled: false,
					autoMemoryEnabled: false,
					precomputeCompactionEnabled: false,
				},
				mcpServers: {
					klaus: {
						type: "sdk",
						name: "klaus",
						instance: bridge.server,
					},
				},
				stderr: (data) => {
					dbg?.("sdk.stderr");
					stderr = `${stderr}${data}`.slice(-4096);
				},
				/** `canUseTool` runs only when a tool call would prompt, so the tool
				 * use id is captured here as well, where Claude reports every call
				 * regardless of how permissions resolve. */
				hooks: {
					PreToolUse: [
						{
							hooks: [
								async (input, toolUseId) => {
									if (input.hook_event_name !== "PreToolUse" || !toolUseId) {
										return { continue: true };
									}
									dbg?.("sdk.hook.preToolUse");
									bridge.allow(
										input.tool_name,
										(input.tool_input ?? {}) as Record<string, unknown>,
										toolUseId,
									);
									return { continue: true };
								},
							],
						},
					],
				},
				canUseTool: async (name, input, options) => {
					dbg?.("sdk.canUseTool");
					bridge.allow(name, input, options.toolUseID);
					return {
						behavior: "allow",
						updatedInput: { ...input, [TOOL_ID_FIELD]: options.toolUseID },
						toolUseID: options.toolUseID,
					};
				},
			},
		});
		dbg?.("sdk.query.returned");
	} catch (error) {
		dbg?.("sdk.query.error", { kind: diagnosticKind(error) });
		await close("Klaus query startup failed.");
		throw error;
	}

	const pump = async (): Promise<void> => {
		dbg?.("sdk.pump.start");
		try {
			for await (const message of queryHandle as Query) {
				dbg?.("sdk.pump.message");
				active.onActivity();
				if (
					(message.type === "assistant" || message.type === "user") &&
					message.uuid
				) {
					messageState.latestPosition = message.uuid;
					dbg?.("sdk.position.candidate");
				}
				const terminal = await handleMessage(
					message,
					blocks,
					active,
					messageState,
					bridge.activeTools(),
					() => {
						if (closed) return;
						closed = true;
						abortController.abort("Klaus stopped max-token recovery.");
						queryHandle?.close();
					},
				);
				if (message.type === "result") {
					turnActive = false;
					reusable =
						!closed && message.subtype === "success" && !message.is_error;
				}
				if (terminal) {
					turnActive = false;
					reusable = false;
					closed = true;
					abortController.abort("Klaus stopped max-token recovery.");
					queryHandle?.close();
					break;
				}
			}
		} catch (error) {
			if (!closed) {
				const detail = redact(stderr.trim(), oauthToken);
				const message = redact(errorText(error), oauthToken);
				dbg?.("sdk.pump.error", { kind: diagnosticKind(error) });
				active.onError(
					new Error(
						nameToolSchemaError(
							detail ? `${message}: ${detail}` : message,
							bridge.activeTools(),
						),
					),
				);
			}
		} finally {
			dbg?.("sdk.pump.finally", { closed });
			turnActive = false;
			reusable = false;
			if (!closed) closed = true;
			await cleanup("Klaus query finished.");
			if (messageState.deferredOutputLimit) {
				active.onResult(messageState.deferredOutputLimit);
			}
		}
	};
	pumpPromise = pump();
	dbg?.("sdk.handle.ready");
	return {
		get bridge() {
			return bridge;
		},
		close,
		isIdle: () => reusable && !closed && !turnActive,
		/** Claude caches a server's tool list for as long as that server stays
		 * registered, so a swap publishes a fresh bridge under a fresh name and
		 * drops the previous one. */
		setTools: async (tools) => {
			if (!reusable || closed || turnActive) {
				throw new Error("Klaus query cannot change tools right now.");
			}
			dbg?.("sdk.setTools");
			const previous = bridge;
			bridgeGeneration += 1;
			const name = `klaus-${bridgeGeneration}`;
			const next = new ToolBridge(tools, () => active.onActivity());
			await queryHandle?.setMcpServers({
				[name]: { type: "sdk", name, instance: next.server },
			});
			bridge = next;
			await previous.close("Klaus replaced the tool bridge.");
		},
		send: (turn, turnCallbacks) => {
			if (!reusable || closed || turnActive) {
				throw new Error("Klaus query cannot accept another turn.");
			}
			dbg?.("sdk.send", { imageCount: turn.images?.length ?? 0 });
			for (const image of turn.images ?? []) imageMediaType(image.mimeType);
			active = turnCallbacks;
			startTurn();
			channel.push(userMessage(turn));
		},
	};
}

async function handleMessage(
	message: SDKMessage,
	blocks: Map<number, BlockState>,
	callbacks: KlausQueryCallbacks,
	state: MessageState,
	tools: KlausTool[],
	stopMaxRecovery: () => void,
): Promise<boolean> {
	dbg?.("sdk.handleMessage");
	if (message.type === "system" && message.subtype === "init") {
		dbg?.("sdk.message.init");
		state.sdkSessionId = message.session_id;
		await callbacks.onReady({
			"x-klaus-transport": "claude-agent-sdk",
			"x-klaus-sdk-version": SDK_VERSION,
			"x-klaus-claude-version": CLAUDE_VERSION,
			"x-klaus-session-id": message.session_id,
		});
		return false;
	}
	if (message.type === "assistant") {
		dbg?.("sdk.message.assistant");
		return false;
	}
	if (message.type === "system" && message.subtype === "mirror_error") {
		dbg?.("sdk.message.mirror_error", {
			kind: diagnosticKind(message.error),
		});
		callbacks.onError(
			new Error(`Klaus session mirror failed: ${message.error}`),
		);
		return false;
	}
	if (message.type === "system" && message.subtype === "api_retry") {
		dbg?.("sdk.message.apiRetry", { attempt: message.attempt });
		if (message.attempt > 1) {
			callbacks.onNotice(
				`Klaus is retrying the Claude request (attempt ${message.attempt} of ${message.max_retries}).`,
			);
		}
		return false;
	}
	if (message.type === "system" && message.subtype === "permission_denied") {
		dbg?.("sdk.message.permissionDenied");
		callbacks.onNotice(
			`Claude denied the ${message.tool_name} tool call without asking Klaus: ${message.message}`,
		);
		return false;
	}
	if (message.type === "system" && message.subtype === "worker_shutting_down") {
		dbg?.("sdk.message.workerShuttingDown");
		callbacks.onError(
			new Error(`Klaus child shut down before finishing: ${message.reason}`),
		);
		return false;
	}
	if (
		message.type === "system" &&
		message.subtype === "informational" &&
		message.level === "warning"
	) {
		dbg?.("sdk.message.informational");
		callbacks.onNotice(`Klaus child warning: ${message.content}`);
		return false;
	}
	if (
		message.type === "rate_limit_event" &&
		message.rate_limit_info.status === "rejected"
	) {
		dbg?.("sdk.message.rateLimitRejected");
		callbacks.onError(
			new Error("rate limit: Claude subscription limit rejected the request."),
		);
		return false;
	}
	if (message.type === "result") {
		dbg?.("sdk.message.result");
		if (message.subtype !== "success" || message.is_error) {
			const detail =
				message.subtype === "success"
					? message.result
					: message.errors.join("\n");
			dbg?.("sdk.message.resultError", { kind: diagnosticKind(detail) });
			if (isOutputLimit(detail)) {
				callbacks.onResult({
					usage: remainingUsage(message.usage, state.reportedUsage),
					responseId: message.uuid,
					responseModel: state.responseModel,
					position: state.latestPosition ?? message.uuid,
					sdkSessionId: message.session_id,
					stopReason: "length",
				});
				return false;
			}
			callbacks.onError(
				new Error(
					normalizeProviderError(
						detail || `Claude ended with ${message.subtype}.`,
						tools,
					),
				),
			);
			return false;
		}
		callbacks.onResult({
			usage: remainingUsage(message.usage, state.reportedUsage),
			responseId: message.uuid,
			responseModel: state.responseModel,
			position: state.latestPosition ?? message.uuid,
			sdkSessionId: message.session_id,
			stopReason: message.stop_reason === "max_tokens" ? "length" : "stop",
		});
		return false;
	}
	if (message.type !== "stream_event") return false;
	const event = message.event;
	dbg?.("sdk.streamEvent", {
		index: "index" in event ? event.index : undefined,
	});
	if (event.type === "message_start") {
		state.responseId = event.message.id;
		state.responseModel = event.message.model;
		state.usage.input = event.message.usage.input_tokens;
		state.usage.output = event.message.usage.output_tokens;
		state.usage.cacheRead = event.message.usage.cache_read_input_tokens ?? 0;
		state.usage.cacheWrite =
			event.message.usage.cache_creation_input_tokens ?? 0;
		return false;
	}
	if (event.type === "content_block_start") {
		const block = event.content_block;
		if (block.type === "text") {
			blocks.set(event.index, { type: "text", text: "" });
			callbacks.onContent({ type: "text-start", index: event.index });
		} else if (block.type === "thinking") {
			blocks.set(event.index, { type: "thinking", text: "" });
			callbacks.onContent({ type: "thinking-start", index: event.index });
		} else if (block.type === "redacted_thinking") {
			blocks.set(event.index, {
				type: "thinking",
				text: "",
				signature: block.data,
				redacted: true,
			});
			callbacks.onContent({ type: "thinking-start", index: event.index });
		} else if (block.type === "tool_use") {
			const name = originalToolName(block.name);
			blocks.set(event.index, {
				type: "tool",
				id: block.id,
				name,
				text: "",
			});
			callbacks.onContent({
				type: "tool-start",
				index: event.index,
				id: block.id,
				name,
			});
		}
		return false;
	}
	if (event.type === "content_block_delta") {
		const block = blocks.get(event.index);
		if (!block) return false;
		if (event.delta.type === "text_delta" && block.type === "text") {
			block.text += event.delta.text;
			callbacks.onContent({
				type: "text-delta",
				index: event.index,
				delta: event.delta.text,
			});
		} else if (
			event.delta.type === "thinking_delta" &&
			block.type === "thinking"
		) {
			block.text += event.delta.thinking;
			callbacks.onContent({
				type: "thinking-delta",
				index: event.index,
				delta: event.delta.thinking,
			});
		} else if (
			event.delta.type === "signature_delta" &&
			block.type === "thinking"
		) {
			block.signature = `${block.signature ?? ""}${event.delta.signature}`;
		} else if (
			event.delta.type === "input_json_delta" &&
			block.type === "tool"
		) {
			block.text += event.delta.partial_json;
			callbacks.onContent({
				type: "tool-delta",
				index: event.index,
				delta: event.delta.partial_json,
			});
		}
		return false;
	}
	if (event.type === "content_block_stop") {
		const block = blocks.get(event.index);
		dbg?.("sdk.contentBlock.stop", { index: event.index });
		if (!block) return false;
		blocks.delete(event.index);
		if (block.type === "text") {
			callbacks.onContent({
				type: "text-end",
				index: event.index,
				text: block.text,
			});
		} else if (block.type === "thinking") {
			callbacks.onContent({
				type: "thinking-end",
				index: event.index,
				thinking: block.text,
				signature: block.signature,
				redacted: block.redacted,
			});
		} else {
			let parsed: unknown;
			try {
				parsed = JSON.parse(block.text || "{}");
			} catch {
				throw new Error(
					`Claude emitted malformed JSON for tool ${block.name ?? "unknown"}.`,
				);
			}
			if (
				typeof parsed !== "object" ||
				parsed === null ||
				Array.isArray(parsed)
			) {
				throw new Error(
					`Claude emitted non-object arguments for tool ${block.name ?? "unknown"}.`,
				);
			}
			const arguments_ = parsed as JsonObject;
			callbacks.onContent({
				type: "tool-end",
				index: event.index,
				id: block.id ?? "",
				name: block.name ?? "",
				arguments: arguments_,
			});
		}
		return false;
	}
	if (event.type === "message_delta") {
		dbg?.("sdk.messageDelta");
		state.usage.output = event.usage.output_tokens;
		if (event.delta.stop_reason === "tool_use") {
			callbacks.onToolBoundary({
				usage: { ...state.usage },
				responseId: state.responseId,
				responseModel: state.responseModel,
			});
			addUsage(state.reportedUsage, state.usage);
		} else if (event.delta.stop_reason === "max_tokens") {
			stopMaxRecovery();
			state.deferredOutputLimit = {
				usage: { ...state.usage },
				responseId: state.responseId,
				responseModel: state.responseModel,
				position: undefined,
				sdkSessionId: state.sdkSessionId,
				stopReason: "length",
			};
			return true;
		}
	}
	return false;
}