import { spawn } from "node:child_process"; import { readFile } from "node:fs/promises"; import { homedir } from "node:os"; import { join } from "node:path"; import type { ExtensionAPI } from "@earendil-works/pi-coding-agent"; const PROVIDER_ID = "localai"; const DEFAULT_TIMEOUT_MS = 5_000; const NON_CHAT_MODEL = /(?:^|[-_.])(?:asr|stt|tts|whisper|speech|embed|embedding|embeddings|rerank|bge|e5|flux|sd)(?:$|[-_.])|nomic[-_.]?embed|stable[-_.]?diffusion/i; type FetchLike = typeof fetch; type ProviderSettings = { baseUrl: string; apiKey?: string; compat?: Record; }; type ModelsFile = { providers?: Record; }; type RemoteModel = { id?: unknown; name?: unknown } | string; export type LocalAiModel = { id: string; name: string; reasoning: false; input: ["text"]; cost: { input: 0; output: 0; cacheRead: 0; cacheWrite: 0 }; contextWindow: 128_000; maxTokens: 16_384; compat?: Record; }; function agentDirectory() { return process.env.PI_CODING_AGENT_DIR ?? join(homedir(), ".pi", "agent"); } function runShellCommand(command: string, timeoutMs: number, signal?: AbortSignal) { return new Promise((resolve, reject) => { const child = spawn("/bin/sh", ["-lc", command], { detached: true, stdio: ["ignore", "pipe", "pipe"], }); let stdout = ""; let settled = false; let forceKill: ReturnType | undefined; const killGroup = (signalName: NodeJS.Signals) => { if (!child.pid) return; try { process.kill(-child.pid, signalName); } catch {} }; const cleanup = () => { clearTimeout(timeout); signal?.removeEventListener("abort", onAbort); }; const finish = (error?: Error, value?: string) => { if (settled) return; settled = true; cleanup(); child.stdout.destroy(); child.stderr.destroy(); child.unref(); if (error) reject(error); else resolve(value ?? ""); }; const terminate = (error: Error) => { killGroup("SIGTERM"); forceKill = setTimeout(() => killGroup("SIGKILL"), 250); forceKill.unref(); finish(error); }; const onAbort = () => terminate( signal?.reason instanceof Error ? signal.reason : new DOMException("Cancelled", "AbortError"), ); const timeout = setTimeout( () => terminate(new DOMException("Command timed out", "TimeoutError")), timeoutMs, ); timeout.unref(); if (signal?.aborted) { onAbort(); return; } signal?.addEventListener("abort", onAbort, { once: true }); child.stdout.setEncoding("utf8"); child.stdout.on("data", (chunk: string) => { stdout += chunk; if (stdout.length > 64 * 1024) { terminate(new Error("Command output exceeded 64KB")); } }); child.on("error", (error) => finish(error)); child.on("close", (code) => { if (code === 0) finish(undefined, stdout.trim()); else finish(new Error(`Command exited with status ${code ?? "unknown"}`)); }); }); } export async function resolveConfigValue( value: string | undefined, timeoutMs = DEFAULT_TIMEOUT_MS, signal?: AbortSignal, ) { if (!value) return undefined; if (value.startsWith("!")) { const command = value.slice(1).trim(); if (!command) return undefined; return (await runShellCommand(command, timeoutMs, signal)) || undefined; } const dollar = "\u0000LOCALAI_DOLLAR\u0000"; const bang = "\u0000LOCALAI_BANG\u0000"; let resolved = value.replace(/\$\$/g, dollar).replace(/\$!/g, bang); let missing = false; resolved = resolved.replace( /\$\{([A-Za-z_][A-Za-z0-9_]*)\}|\$([A-Za-z_][A-Za-z0-9_]*)/g, (_match, braced: string | undefined, plain: string | undefined) => { const environmentValue = process.env[braced ?? plain ?? ""]; if (!environmentValue) missing = true; return environmentValue ?? ""; }, ); if (missing) return undefined; return resolved.replaceAll(dollar, "$").replaceAll(bang, "!"); } export function modelsEndpoint(baseUrl: string) { return `${baseUrl.replace(/\/+$/, "")}/models`; } export function mapLocalAiModels( entries: RemoteModel[], compat?: Record, ): LocalAiModel[] { const seen = new Set(); const models: LocalAiModel[] = []; for (const entry of entries) { const id = (typeof entry === "string" ? entry : entry?.id)?.toString().trim(); if (!id || seen.has(id) || NON_CHAT_MODEL.test(id)) continue; seen.add(id); const remoteName = typeof entry === "object" && typeof entry.name === "string" ? entry.name.trim() : ""; models.push({ id, name: remoteName || `${id} (LocalAI)`, reasoning: false, input: ["text"], cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, contextWindow: 128_000, maxTokens: 16_384, ...(compat ? { compat: { ...compat } } : {}), }); } return models.sort((left, right) => left.name.localeCompare(right.name)); } export async function discoverLocalAiModels( settings: ProviderSettings, options: { fetch?: FetchLike; signal?: AbortSignal; timeoutMs?: number } = {}, ) { const timeoutMs = options.timeoutMs ?? DEFAULT_TIMEOUT_MS; const fetchImpl = options.fetch ?? fetch; const timeout = AbortSignal.timeout(timeoutMs); const signal = options.signal ? AbortSignal.any([options.signal, timeout]) : timeout; const apiKey = await resolveConfigValue(settings.apiKey, timeoutMs, signal); const response = await fetchImpl(modelsEndpoint(settings.baseUrl), { headers: apiKey ? { Authorization: `Bearer ${apiKey}` } : undefined, signal, }); if (!response.ok) throw new Error(`LocalAI model discovery failed with HTTP ${response.status}`); const payload = await response.json() as { data?: unknown; models?: unknown }; const entries = Array.isArray(payload.data) ? payload.data : Array.isArray(payload.models) ? payload.models : undefined; if (!entries) throw new Error("LocalAI model discovery returned a malformed payload"); return mapLocalAiModels(entries as RemoteModel[], settings.compat); } export function createLocalAiModelsExtension(options: { configPath?: string; fetch?: FetchLike; timeoutMs?: number; } = {}) { return async function localAiModelsExtension(pi: ExtensionAPI) { const configPath = options.configPath ?? join(agentDirectory(), "models.json"); let settings: ProviderSettings; try { const config = JSON.parse(await readFile(configPath, "utf8")) as ModelsFile; const provider = config.providers?.[PROVIDER_ID]; if (!provider || typeof provider.baseUrl !== "string" || !provider.baseUrl.trim()) { console.warn("LocalAI provider configuration is missing; continuing without LocalAI models."); return; } settings = { baseUrl: provider.baseUrl, apiKey: typeof provider.apiKey === "string" ? provider.apiKey : undefined, compat: provider.compat && typeof provider.compat === "object" ? { ...(provider.compat as Record) } : undefined, }; } catch { console.warn("LocalAI provider configuration could not be read; continuing without LocalAI models."); return; } let startupFailed = false; let initialModels: LocalAiModel[] = []; try { initialModels = await discoverLocalAiModels(settings, { fetch: options.fetch, timeoutMs: options.timeoutMs, }); } catch { startupFailed = true; console.warn("LocalAI model discovery unavailable; continuing without LocalAI models."); } pi.registerProvider(PROVIDER_ID, { name: "LocalAI", baseUrl: settings.baseUrl, apiKey: settings.apiKey, api: "openai-completions", models: initialModels, async refreshModels({ signal }: { signal: AbortSignal }) { try { return await discoverLocalAiModels(settings, { fetch: options.fetch, signal, timeoutMs: options.timeoutMs, }); } catch { console.warn("LocalAI model refresh unavailable; continuing without LocalAI models."); return []; } }, }); if (startupFailed) { pi.on("session_start", (_event, ctx) => { if (ctx.hasUI) ctx.ui.notify("LocalAI unavailable; continuing without LocalAI models.", "warning"); }); } }; } export default createLocalAiModelsExtension();