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}