Parent directory

index.ts

20682 bytes
  1/**
  2 * llama.cpp provider for pi.
  3 *
  4 * Auto-discovers models from a running `llama-server` and
  5 * registers them under the `llama-cpp` provider.
  6 *
  7 * Usage: `pi install github.com/huggingface/pi-llama`
  8 */
  9
 10import { CONFIG_DIR_NAME, getAgentDir, type ExtensionAPI } from "@earendil-works/pi-coding-agent";
 11import { existsSync, readFileSync } from "node:fs";
 12import { join } from "node:path";
 13import { Type } from "typebox";
 14import { Compile } from "typebox/compile";
 15import { Loader, truncateToWidth, visibleWidth } from "@earendil-works/pi-tui";
 16
 17const PROVIDER_ID = "llama-cpp";
 18const DEFAULT_BASE_URL = "http://localhost:8080/v1";
 19// Fallback for /v1/models entries missing meta.n_ctx.
 20const DEFAULT_CONTEXT_WINDOW = 8192;
 21// llama.cpp has no output-token cap (no endpoint reports one; generation is only
 22// bounded by the context window), so use Pi's own default for models that omit
 23// maxTokens (see model-registry.ts parseModels).
 24const DEFAULT_MAX_TOKENS = 16384;
 25const PROPS_TIMEOUT_MS = 120_000;
 26const CONFIG_FILE_NAME = "pi-llama.json";
 27
 28type LlamaConfig = {
 29	autoloadOnSelect?: boolean;
 30	baseUrl?: string;
 31};
 32
 33const DEFAULT_CONFIG: Required<LlamaConfig> = {
 34	autoloadOnSelect: false,
 35	baseUrl: DEFAULT_BASE_URL,
 36};
 37
 38function normalizeBaseUrl(url: string): string {
 39	return url.replace(/\/+$/, "");
 40}
 41
 42function readConfig(path: string): LlamaConfig {
 43	if (!existsSync(path)) {
 44		return {};
 45	}
 46
 47	try {
 48		const value: unknown = JSON.parse(readFileSync(path, "utf8"));
 49		if (!value || typeof value !== "object" || Array.isArray(value)) {
 50			throw new Error("expected a JSON object");
 51		}
 52
 53		const config = value as Record<string, unknown>;
 54		const result: LlamaConfig = {};
 55		if (typeof config.autoloadOnSelect === "boolean") {
 56			result.autoloadOnSelect = config.autoloadOnSelect;
 57		}
 58		if (typeof config.baseUrl === "string" && config.baseUrl.trim()) {
 59			result.baseUrl = normalizeBaseUrl(config.baseUrl.trim());
 60		}
 61		return result;
 62	} catch (error) {
 63		console.warn(`[llama-cpp] failed to read config ${path}: ${(error as Error).message}`);
 64		return {};
 65	}
 66}
 67
 68function loadConfig(cwd: string, projectIsTrusted: boolean): Required<LlamaConfig> {
 69	const globalConfig = readConfig(join(getAgentDir(), "extensions", CONFIG_FILE_NAME));
 70	const projectConfig = projectIsTrusted
 71		? readConfig(join(cwd, CONFIG_DIR_NAME, CONFIG_FILE_NAME))
 72		: {};
 73	const config = { ...DEFAULT_CONFIG, ...globalConfig, ...projectConfig };
 74
 75	return {
 76		...config,
 77		baseUrl: normalizeBaseUrl(process.env.LLAMA_BASE_URL ?? config.baseUrl),
 78	};
 79}
 80
 81const ModelsResponseSchema = Type.Object({
 82	data: Type.Optional(
 83		Type.Array(
 84			Type.Object({
 85				id: Type.String(),
 86				aliases: Type.Optional(Type.Array(Type.String())),
 87				status: Type.Optional(
 88					Type.Object({
 89						value: Type.Optional(
 90							Type.Union([
 91								Type.Literal("unloaded"),
 92								Type.Literal("loading"),
 93								Type.Literal("loaded"),
 94								Type.Literal("sleeping"),
 95								Type.Literal("unknown"),
 96							]),
 97						),
 98					}),
 99				),
100				architecture: Type.Optional(
101					Type.Object({
102						input_modalities: Type.Optional(Type.Array(Type.String())),
103					}),
104				),
105				meta: Type.Optional(
106					Type.Object({
107						n_ctx: Type.Optional(Type.Number()),
108						n_params: Type.Optional(Type.Number()),
109					}),
110				),
111			}),
112		),
113	),
114});
115
116const validateModelsResponse = Compile(ModelsResponseSchema);
117
118const PropsResponseSchema = Type.Object({
119	default_generation_settings: Type.Optional(
120		Type.Object({
121			n_ctx: Type.Optional(Type.Number()),
122		}),
123	),
124	chat_template: Type.Optional(Type.String()),
125	build_info: Type.Optional(Type.String()),
126});
127
128const validatePropsResponse = Compile(PropsResponseSchema);
129
130// SSE event types for model loading progress
131type ApiModelLoadStage = "text_model" | "spec_model" | "mmproj_model";
132
133type ApiModelsSseProgress = {
134	stages: ApiModelLoadStage[];
135	current: ApiModelLoadStage;
136	value: number;
137};
138
139type ApiModelsSseData = {
140	status: string;
141	progress?: ApiModelsSseProgress;
142	exit_code?: number;
143};
144
145type ApiModelsSseEvent = {
146	model: string;
147	event: string;
148	data: ApiModelsSseData;
149};
150
151const MODEL_LOAD_STAGE_LABELS: Record<ApiModelLoadStage, string> = {
152	text_model: "Loading weights",
153	spec_model: "Loading draft",
154	mmproj_model: "Loading projector",
155};
156
157type LlamaModel = NonNullable<Parameters<ExtensionAPI["registerProvider"]>[1]["models"]>[number];
158type ExtensionCtx = Parameters<Parameters<ExtensionAPI["on"]>[1]>[1];
159
160// llama.cpp template thinking is boolean, so expose Pi's default off/medium toggle only.
161const TEMPLATE_THINKING_LEVEL_MAP = {
162	minimal: null,
163	low: null,
164	high: null,
165	xhigh: null,
166} satisfies NonNullable<LlamaModel["thinkingLevelMap"]>;
167
168// Minimal shape needed to update both registered models and Pi's active model snapshot.
169type MutableModelMetadata = {
170	reasoning: boolean;
171	thinkingLevelMap?: LlamaModel["thinkingLevelMap"];
172	compat?: LlamaModel["compat"];
173	contextWindow: number;
174	maxTokens: number;
175};
176
177// Mark a model as using llama.cpp's chat_template_kwargs.enable_thinking control.
178function applyTemplateThinkingSupport(model: MutableModelMetadata): void {
179	model.reasoning = true;
180	model.thinkingLevelMap = TEMPLATE_THINKING_LEVEL_MAP;
181	model.compat = {
182		...model.compat,
183		// Despite the Pi enum name, this sends llama.cpp's generic
184		// chat_template_kwargs.enable_thinking payload, not a Qwen-only option.
185		thinkingFormat: "qwen-chat-template",
186	};
187}
188
189// Pi invalidates a captured ctx when the session is replaced (e.g. new_session in
190// RPC mode). Any later ctx access then throws this error. Background work started
191// before the replacement should treat it as "session gone" and stop quietly.
192function isStaleContextError(error: unknown): boolean {
193	return error instanceof Error && error.message.includes("stale after session replacement");
194}
195
196export default async function (pi: ExtensionAPI) {
197	let currentModels: LlamaModel[] = [];
198	let config = loadConfig(process.cwd(), false);
199	let autoloadOnSelect = config.autoloadOnSelect;
200
201	pi.registerCommand("llama-version", {
202		description: "Get build info of llama.cpp server",
203		handler: async (_args, ctx) => {
204			const response = await fetch(`${baseUrl.replace(/\/v1$/, "")}/props`);
205			if (!response.ok) {
206				ctx.ui.notify(`[llama-cpp] /props returned ${response.status}`, "error");
207				return;
208			}
209
210			const data: unknown = await response.json();
211			if (!validatePropsResponse.Check(data)) {
212				const errors = [...validatePropsResponse.Errors(data)]
213					.map((e) => `${"path" in e ? e.path : ""} ${e.message}`)
214					.join("; ");
215				ctx.ui.notify(`[llama-cpp] invalid /props response: ${errors}`, "error");
216				return;
217			}
218
219			const match = data.build_info?.match(/^b([a-zA-Z0-9]+)-([a-zA-Z0-9]+)$/);
220
221			if (match && match.length === 3) {
222				ctx.ui.notify(`Build number: ${match[1]}, Commit hash: ${match[2]}`, "info");
223			} else {
224				ctx.ui.notify(`Malformed build info: ${data.build_info}`, "warning");
225			}
226		},
227	});
228
229	let baseUrl = config.baseUrl;
230	const apiKey = process.env.LLAMA_API_KEY ?? "no-key";
231
232	async function refreshProvider(): Promise<void> {
233		try {
234			const response = await fetch(`${baseUrl}/models`);
235			if (!response.ok) {
236				console.warn(`[llama-cpp] ${baseUrl}/models returned ${response.status}`);
237				return;
238			}
239
240			const payload: unknown = await response.json();
241			if (!validateModelsResponse.Check(payload)) {
242				const errors = [...validateModelsResponse.Errors(payload)]
243					.map((e) => `${"path" in e ? e.path : ""} ${e.message}`)
244					.join("; ");
245				console.warn(`[llama-cpp] invalid /models response: ${errors}`);
246				return;
247			}
248
249			const previousById = new Map(currentModels.map((m) => [m.id, m]));
250
251			currentModels = (payload.data ?? []).map((model) => {
252				const previous = previousById.get(model.id);
253				const isLoaded = model.status?.value === "loaded";
254				const modalities = model.architecture?.input_modalities ?? ["text"];
255				const input = modalities.filter(
256					(m): m is "text" | "image" => m === "text" || m === "image",
257				);
258				const suffixes: string[] = [];
259				if (input.includes("image")) {
260					suffixes.push("(image)");
261				}
262				if (isLoaded) {
263					suffixes.push("(loaded)");
264				}
265				const contextWindow =
266					model.meta?.n_ctx ?? previous?.contextWindow ?? DEFAULT_CONTEXT_WINDOW;
267				const displayName = model.aliases?.[0] || model.id;
268				return {
269					id: model.id,
270					name: suffixes.length > 0 ? `${displayName} ${suffixes.join(" ")}` : displayName,
271					// /v1/models does not include /props-discovered capabilities, so preserve
272					// template thinking metadata across refreshes.
273					reasoning: previous?.reasoning ?? false,
274					thinkingLevelMap: previous?.thinkingLevelMap,
275					input,
276					cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 },
277					contextWindow,
278					maxTokens: Math.min(DEFAULT_MAX_TOKENS, contextWindow),
279					compat: previous?.compat,
280					status: model.status,
281				} as LlamaModel;
282			});
283
284			if (currentModels.length === 0) {
285				console.warn(`[llama-cpp] no models returned from ${baseUrl}/models`);
286				return;
287			}
288
289			// Track which model is currently loaded on the server
290			const loadedModel = currentModels.find((m) => m.status?.value === "loaded");
291			currentlyLoadedModel = loadedModel?.id ?? null;
292
293			pi.registerProvider(PROVIDER_ID, {
294				name: "llama.cpp",
295				baseUrl,
296				apiKey,
297				api: "openai-completions",
298				models: currentModels,
299			});
300		} catch (error) {
301			console.warn(`[llama-cpp] failed to reach ${baseUrl}/models: ${(error as Error).message}`);
302		}
303	}
304
305	const discoveredMetadata = new Set<string>();
306	const pendingMetadata = new Set<string>();
307	let currentlyLoadedModel: string | null = null;
308	let statusTimeout: ReturnType<typeof setTimeout> | undefined;
309	let sseAbortController: AbortController | null = null;
310	let propsAbortController: AbortController | null = null;
311
312	function clearFooterStatusTimeout(): void {
313		if (statusTimeout !== undefined) {
314			clearTimeout(statusTimeout);
315			statusTimeout = undefined;
316		}
317	}
318
319	// Connect to SSE stream for model loading progress
320	async function connectToLoadingProgress(
321		modelId: string,
322		ctx: ExtensionCtx,
323		loader: Loader,
324	): Promise<void> {
325		// Close any existing SSE connection
326		if (sseAbortController) {
327			sseAbortController.abort();
328			sseAbortController = null;
329		}
330
331		sseAbortController = new AbortController();
332		const signal = sseAbortController.signal;
333
334		try {
335			const response = await fetch(`${baseUrl.replace(/\/v1$/, "")}/models/sse`, { signal });
336
337			if (!response.ok) {
338				if (response.status !== 404) {
339					ctx?.ui.notify(`[llama-cpp] loading progress ${response.status})`, "warning");
340				}
341				return;
342			}
343
344			const reader = response.body?.getReader();
345			if (!reader) {
346				return;
347			}
348
349			const decoder = new TextDecoder();
350			let buffer = "";
351
352			while (!signal.aborted) {
353				const { done, value } = await reader.read();
354				if (done) {
355					break;
356				}
357
358				buffer += decoder.decode(value, { stream: true });
359				const events = buffer.split("\n\n");
360				buffer = events.pop() || "";
361
362				for (const event of events) {
363					if (!event) {
364						continue;
365					}
366
367					// Parse SSE record: extract data lines
368					const dataLines = event
369						.split("\n")
370						.filter((line) => line.startsWith("data:"))
371						.map((line) => line.slice(5).trim())
372						.join("\n");
373
374					if (!dataLines) {
375						continue;
376					}
377
378					try {
379						const sseEvent: ApiModelsSseEvent = JSON.parse(dataLines);
380
381						// Process status events for all models to keep discoveredMetadata and
382						// currentlyLoadedModel in sync.
383						if (
384							sseEvent.event === "model_status" ||
385							sseEvent.event === "status_change" ||
386							sseEvent.event === "status_update"
387						) {
388							const status = sseEvent.data.status;
389
390							if (status === "unloaded") {
391								discoveredMetadata.delete(sseEvent.model);
392								if (currentlyLoadedModel === sseEvent.model) {
393									currentlyLoadedModel = null;
394								}
395							}
396							if (status === "loaded") {
397								currentlyLoadedModel = sseEvent.model;
398							}
399						}
400
401						// Progress UI is only for the model we're actively loading
402						if (sseEvent.model === modelId) {
403							const currentModel = currentModels.find((m) => m.id === sseEvent.model);
404							const displayName = currentModel?.name.split(" ")[0] || sseEvent.model;
405							const progress = sseEvent.data.progress;
406
407							if (sseEvent.data.exit_code && sseEvent.data.exit_code !== 0) {
408								ctx?.ui.setWidget(PROVIDER_ID, [
409									ctx?.ui.theme.fg("error", "[llama.cpp] ") +
410										ctx?.ui.theme.fg(
411											"text",
412											`${displayName}:  failed (exit ${sseEvent.data.exit_code})`,
413										),
414								]);
415								sseAbortController?.abort();
416								return;
417							}
418
419							if (sseEvent.data.status === "loading" && progress) {
420								const stageLabel = progress.current
421									? MODEL_LOAD_STAGE_LABELS[progress.current] || progress.current
422									: "Loading";
423								const progressPercent = Math.round(progress.value * 100);
424								loader?.setMessage(`${displayName}: ${stageLabel} (${progressPercent}%)`);
425							}
426						}
427					} catch {
428						// Ignore parse errors
429					}
430				}
431			}
432		} catch (error) {
433			// Suppress errors from intentionally-aborted SSE connections and from a
434			// stale ctx (session was replaced while streaming).
435			const msg = (error as Error).message;
436			if (
437				isStaleContextError(error) ||
438				(signal.aborted && (error instanceof DOMException || msg === "terminated"))
439			) {
440				return;
441			}
442			ctx?.ui.notify(`[llama-cpp] SSE error: ${msg}`, "warning");
443		} finally {
444			sseAbortController = null;
445		}
446	}
447
448	async function discoverModelMetadata(
449		modelId: string,
450		ctx?: ExtensionCtx,
451		autoload = true,
452		timeoutMs = PROPS_TIMEOUT_MS,
453		selectedModel?: MutableModelMetadata,
454	): Promise<void> {
455		const model = currentModels.find((m) => m.id === modelId);
456		if (!model) {
457			return;
458		}
459		const displayName = model.name.split(" ")[0];
460		// Use tracked state instead of stale currentModels status.
461		const isLoaded = currentlyLoadedModel === modelId;
462
463		// llama.cpp rejects /props?autoload=false for unloaded models.
464		if (!autoload && !isLoaded) {
465			return;
466		}
467
468		if (discoveredMetadata.has(modelId)) {
469			// If discovered but no longer loaded, clear cache and fall through to reload.
470			if (!isLoaded) {
471				discoveredMetadata.delete(modelId);
472			} else {
473				// Copy cached metadata into the selected model snapshot.
474				if (selectedModel) {
475					selectedModel.contextWindow = model.contextWindow;
476					selectedModel.maxTokens = model.maxTokens;
477					if (model.reasoning) {
478						selectedModel.reasoning = model.reasoning;
479						selectedModel.thinkingLevelMap = model.thinkingLevelMap;
480						selectedModel.compat = model.compat;
481					}
482				}
483				return;
484			}
485		}
486		if (pendingMetadata.has(modelId)) {
487			return;
488		}
489
490		pendingMetadata.add(modelId);
491		// Cancel any pending clear timeout from a previous model load.
492		clearFooterStatusTimeout();
493		// Abort any in-flight /props request from a previous model.
494		if (propsAbortController) {
495			propsAbortController.abort();
496		}
497		propsAbortController = new AbortController();
498		const timer = setTimeout(() => propsAbortController.abort(), timeoutMs);
499		const shouldAutoload = autoload && !isLoaded;
500		const propsUrl = `${baseUrl.replace(/\/v1$/, "")}/props?model=${encodeURIComponent(modelId)}&autoload=${shouldAutoload}`;
501		const clearFooterStatusLater = () => {
502			clearFooterStatusTimeout();
503			statusTimeout = setTimeout(() => {
504				statusTimeout = undefined;
505				ctx?.ui.setWidget(PROVIDER_ID, undefined);
506			}, 8000);
507		};
508
509		try {
510			if (shouldAutoload && ctx) {
511				let loader = null;
512				ctx.ui.setWidget(PROVIDER_ID, (ui, theme) => {
513					const prefix = theme.fg("accent", " [llama.cpp]");
514					const prefixWidth = visibleWidth(" [llama.cpp]");
515					loader = new Loader(
516						ui,
517						(s) => theme.fg("accent", s),
518						(t) => theme.fg("text", t),
519						`${displayName}: Loading...`,
520					);
521					return {
522						dispose: () => loader?.stop(),
523						render: (width: number) => {
524							const [_, line] = loader.render(width - prefixWidth);
525							return [prefix + truncateToWidth(line, width - prefixWidth)];
526						},
527					};
528				});
529				// Start SSE connection to monitor loading progress
530				void connectToLoadingProgress(modelId, ctx, loader);
531			}
532
533			const response = await fetch(propsUrl, { signal: propsAbortController.signal });
534			if (!response.ok) {
535				// 500 during autoload is expected when the server cancels a load to start
536				// another model. Suppress the notification for that case.
537				if (!(shouldAutoload && response.status === 500)) {
538					ctx?.ui.notify(`[llama-cpp] /props for ${modelId} returned ${response.status}`, "error");
539				}
540				return;
541			}
542			const data: unknown = await response.json();
543			if (!validatePropsResponse.Check(data)) {
544				const errors = [...validatePropsResponse.Errors(data)]
545					.map((e) => `${"path" in e ? e.path : ""} ${e.message}`)
546					.join("; ");
547				ctx?.ui.notify(`[llama-cpp] invalid /props response for ${modelId}: ${errors}`, "error");
548				return;
549			}
550			const nCtx = data.default_generation_settings?.n_ctx;
551			let updated = false;
552			let loadedFooterStatus = shouldAutoload ? `[llama.cpp] ${displayName} loaded` : undefined;
553			if (typeof nCtx === "number" && nCtx > 0) {
554				model.contextWindow = nCtx;
555				model.maxTokens = Math.min(DEFAULT_MAX_TOKENS, nCtx);
556				loadedFooterStatus = `[llama.cpp] ${displayName} loaded with ctx ${nCtx} tokens`;
557				updated = true;
558			}
559			if (selectedModel) {
560				selectedModel.contextWindow = model.contextWindow;
561				selectedModel.maxTokens = model.maxTokens;
562			}
563			if (data.chat_template?.includes("enable_thinking") === true) {
564				applyTemplateThinkingSupport(model);
565				if (selectedModel) {
566					applyTemplateThinkingSupport(selectedModel);
567					if (pi.getThinkingLevel() === "off") {
568						pi.setThinkingLevel("medium");
569					}
570				}
571				updated = true;
572			}
573			discoveredMetadata.add(modelId);
574			if (shouldAutoload) {
575				currentlyLoadedModel = modelId;
576			}
577			if (loadedFooterStatus && ctx && !isLoaded) {
578				const prefix = ctx.ui.theme.fg("success", "[llama.cpp] ✓");
579				ctx.ui.setWidget(PROVIDER_ID, [
580					prefix +
581						ctx.ui.theme.fg(
582							"text",
583							` ${displayName}: Loaded` + (nCtx ? ` with context ${nCtx} tokens` : ""),
584						),
585				]);
586				clearFooterStatusLater();
587			}
588			if (!updated) {
589				return;
590			}
591			pi.registerProvider(PROVIDER_ID, {
592				name: "llama.cpp",
593				baseUrl,
594				apiKey,
595				api: "openai-completions",
596				models: currentModels,
597			});
598		} catch (error) {
599			const err = error as Error;
600			// Suppress notification for aborted requests (model was switched) and for a
601			// stale ctx (session was replaced while awaiting) — both are expected.
602			if (err.name !== "AbortError" && !isStaleContextError(err)) {
603				ctx?.ui.notify(`[llama-cpp] /props for ${modelId} failed: ${err.message}`, "error");
604			}
605		} finally {
606			clearTimeout(timer);
607			pendingMetadata.delete(modelId);
608			propsAbortController = null;
609			// Stop SSE connection when done
610			if (sseAbortController) {
611				sseAbortController.abort();
612				sseAbortController = null;
613			}
614		}
615	}
616
617	pi.on("session_start", async (_event, ctx) => {
618		config = loadConfig(ctx.cwd, ctx.isProjectTrusted());
619		autoloadOnSelect = config.autoloadOnSelect;
620		if (baseUrl !== config.baseUrl) {
621			baseUrl = config.baseUrl;
622			await refreshProvider();
623		}
624	});
625
626	await refreshProvider();
627
628	pi.on("input", async (event) => {
629		const trimmed = event.text.trim().toLowerCase();
630		if (trimmed === "/model") {
631			await refreshProvider();
632		}
633	});
634
635	pi.on("model_select", (event, ctx) => {
636		if (event.model.provider !== PROVIDER_ID) {
637			return;
638		}
639		void discoverModelMetadata(event.model.id, ctx, autoloadOnSelect, PROPS_TIMEOUT_MS, event.model);
640	});
641
642	// Discover /props for already-active models because re-selecting them does not emit model_select.
643	pi.on("before_provider_request", (event, ctx) => {
644		try {
645			const modelId = (event.payload as { model?: unknown })?.model;
646			if (typeof modelId === "string") {
647				const activeModel =
648					ctx.model?.provider === PROVIDER_ID && ctx.model.id === modelId ? ctx.model : undefined;
649				void discoverModelMetadata(modelId, ctx, true, PROPS_TIMEOUT_MS, activeModel);
650			}
651		} catch (error) {
652			// Session was replaced as the request fired; nothing to discover.
653			if (!isStaleContextError(error)) {
654				throw error;
655			}
656		}
657	});
658
659	pi.on("session_shutdown", () => {
660		clearFooterStatusTimeout();
661		// Stop in-flight /props and SSE so they don't resume against a stale ctx.
662		propsAbortController?.abort();
663		sseAbortController?.abort();
664	});
665}