From bcf729f39ccab078cd533a009c0a73130cf33cc8 Mon Sep 17 00:00:00 2001 From: Alex Blank Date: Thu, 27 Aug 2026 16:26:41 +0200 Subject: [PATCH] feat(pi): discover LocalAI models dynamically --- .../extensions/localai-models/index.test.ts | 367 ++++++++++++++++++ .pi/agent/extensions/localai-models/index.ts | 252 ++++++++++++ 2 files changed, 619 insertions(+) create mode 100644 .pi/agent/extensions/localai-models/index.test.ts create mode 100644 .pi/agent/extensions/localai-models/index.ts diff --git a/.pi/agent/extensions/localai-models/index.test.ts b/.pi/agent/extensions/localai-models/index.test.ts new file mode 100644 index 0000000..2834693 --- /dev/null +++ b/.pi/agent/extensions/localai-models/index.test.ts @@ -0,0 +1,367 @@ +import assert from "node:assert/strict"; +import { spawn } from "node:child_process"; +import { mkdtemp, mkdir, readFile, rm, writeFile } from "node:fs/promises"; +import { createServer } from "node:http"; +import { tmpdir } from "node:os"; +import { dirname, join } from "node:path"; +import { fileURLToPath } from "node:url"; +import test from "node:test"; +import { + createLocalAiModelsExtension, + discoverLocalAiModels, + mapLocalAiModels, + resolveConfigValue, +} from "./index.ts"; + +function response(data: unknown, status = 200) { + return new Response(JSON.stringify(data), { + status, + headers: { "content-type": "application/json" }, + }); +} + +async function configFile() { + const root = await mkdtemp(join(tmpdir(), "localai-models-")); + const path = join(root, "models.json"); + await writeFile(path, JSON.stringify({ + providers: { + localai: { + baseUrl: "http://local.invalid/v1", + apiKey: "local-test", + compat: { + supportsDeveloperRole: false, + supportsReasoningEffort: false, + maxTokensField: "max_tokens", + }, + }, + }, + })); + return { root, path }; +} + +function abortAwareFetch(): typeof fetch { + return async (_url, init) => new Promise((_resolve, reject) => { + const signal = init?.signal; + const abort = () => reject(signal?.reason ?? new DOMException("Aborted", "AbortError")); + if (signal?.aborted) abort(); + else signal?.addEventListener("abort", abort, { once: true }); + }); +} + +async function runPiCatalog(agentDir: string) { + const sessionDir = join(agentDir, "sessions"); + await mkdir(sessionDir, { recursive: true }); + const child = spawn("pi", ["--mode", "rpc", "--session-dir", sessionDir], { + cwd: agentDir, + detached: true, + env: { ...process.env, PI_CODING_AGENT_DIR: agentDir }, + stdio: ["pipe", "pipe", "pipe"], + }); + try { + return await new Promise((resolve, reject) => { + let buffer = ""; + const timeout = setTimeout(() => reject(new Error("Pi RPC catalog probe timed out")), 15_000); + child.stdout.setEncoding("utf8"); + child.stdout.on("data", (chunk: string) => { + buffer += chunk; + while (buffer.includes("\n")) { + const newline = buffer.indexOf("\n"); + const line = buffer.slice(0, newline); + buffer = buffer.slice(newline + 1); + if (!line.trim()) continue; + const frame = JSON.parse(line); + if (frame.type === "response" && frame.id === "catalog-test") { + clearTimeout(timeout); + if (!frame.success) reject(new Error(frame.error ?? "Pi RPC catalog probe failed")); + else resolve(frame.data?.models ?? []); + } + } + }); + child.once("error", reject); + child.once("exit", (code) => reject(new Error(`Pi exited before catalog response (${code})`))); + child.stdin.write(`${JSON.stringify({ type: "get_available_models", id: "catalog-test" })}\n`); + }); + } finally { + if (child.pid) { + try { process.kill(-child.pid, "SIGTERM"); } catch {} + } + } +} + +test("maps unique chat models, stamps compatibility, and filters obvious non-chat families", () => { + const compat = { supportsDeveloperRole: false, maxTokensField: "max_tokens" }; + const models = mapLocalAiModels([ + { id: "qwen3.8-27b-q4" }, + { id: "qwen3.8-27b-heretic-abliterated-uncensored" }, + { id: "qwen3.8-27b-q4" }, + { id: "qwen3-asr-1.7b" }, + { id: "whisper-large-v3" }, + { id: "bge-m3" }, + { id: "nomic-embed-text" }, + { id: "flux.2-klein-9b" }, + { id: "sd-3.5-large-ggml" }, + ], compat); + assert.deepEqual(models.map((model) => model.id), [ + "qwen3.8-27b-heretic-abliterated-uncensored", + "qwen3.8-27b-q4", + ]); + assert.deepEqual(models[0].compat, compat); +}); + +test("resolves literal, interpolated, escaped, and command-backed config values", async () => { + const previous = process.env.LOCALAI_TEST_KEY; + const previousEmpty = process.env.LOCALAI_EMPTY_KEY; + process.env.LOCALAI_TEST_KEY = "environment-key"; + process.env.LOCALAI_EMPTY_KEY = ""; + try { + assert.equal(await resolveConfigValue("literal-key"), "literal-key"); + assert.equal(await resolveConfigValue("$LOCALAI_TEST_KEY"), "environment-key"); + assert.equal(await resolveConfigValue("prefix-${LOCALAI_TEST_KEY}"), "prefix-environment-key"); + assert.equal(await resolveConfigValue("$$money-$!bang"), "$money-!bang"); + assert.equal(await resolveConfigValue("!printf command-key"), "command-key"); + assert.equal(await resolveConfigValue("$MISSING_LOCALAI_KEY"), undefined); + assert.equal(await resolveConfigValue("$LOCALAI_EMPTY_KEY"), undefined); + assert.equal(await resolveConfigValue("prefix-${LOCALAI_EMPTY_KEY}"), undefined); + } finally { + if (previous === undefined) delete process.env.LOCALAI_TEST_KEY; + else process.env.LOCALAI_TEST_KEY = previous; + if (previousEmpty === undefined) delete process.env.LOCALAI_EMPTY_KEY; + else process.env.LOCALAI_EMPTY_KEY = previousEmpty; + } +}); + +test("aborts command-backed auth resolution with the caller signal", async () => { + const root = await mkdtemp(join(tmpdir(), "localai-command-abort-")); + const pidFile = join(root, "child.pid"); + const controller = new AbortController(); + const startedAt = Date.now(); + try { + const resolving = resolveConfigValue( + `!sleep 30 & child=$!; printf '%s' "$child" > '${pidFile}'; wait`, + 20_000, + controller.signal, + ); + let childPid = 0; + for (let attempt = 0; attempt < 100 && !childPid; attempt += 1) { + try { + childPid = Number((await readFile(pidFile, "utf8")).trim()); + } catch {} + if (!childPid) await new Promise((resolve) => setTimeout(resolve, 5)); + } + assert.ok(childPid > 0, "descendant pid was recorded"); + controller.abort(new DOMException("Cancelled", "AbortError")); + await assert.rejects(resolving, /abort|cancel/i); + for (let attempt = 0; attempt < 100; attempt += 1) { + try { + process.kill(childPid, 0); + } catch { + childPid = 0; + break; + } + await new Promise((resolve) => setTimeout(resolve, 5)); + } + assert.equal(childPid, 0, "descendant process was terminated"); + assert.ok(Date.now() - startedAt < 1_500, "abort settled promptly"); + } finally { + await rm(root, { recursive: true, force: true }); + } +}); + +test("discovers changing payload variants without caching and sends auth only when resolved", async () => { + let call = 0; + const seenHeaders: Array = []; + const fetchImpl: typeof fetch = async (_url, init) => { + call += 1; + seenHeaders.push(init?.headers); + return call === 1 + ? response({ data: [{ id: "qwen3.8-27b-q4" }] }) + : response({ models: [{ id: "new-chat-model" }] }); + }; + const first = await discoverLocalAiModels( + { baseUrl: "http://local.invalid/v1", apiKey: "local" }, + { fetch: fetchImpl }, + ); + const second = await discoverLocalAiModels( + { baseUrl: "http://local.invalid/v1" }, + { fetch: fetchImpl }, + ); + assert.deepEqual(first.map((model) => model.id), ["qwen3.8-27b-q4"]); + assert.deepEqual(second.map((model) => model.id), ["new-chat-model"]); + assert.deepEqual(seenHeaders[0], { Authorization: "Bearer local" }); + assert.equal(seenHeaders[1], undefined); +}); + +test("rejects malformed, failed, timed-out, and externally aborted discovery", async () => { + const settings = { baseUrl: "http://local.invalid/v1" }; + await assert.rejects( + discoverLocalAiModels(settings, { fetch: async () => response({ object: "list" }) }), + /malformed payload/, + ); + await assert.rejects( + discoverLocalAiModels(settings, { fetch: async () => response({}, 503) }), + /HTTP 503/, + ); + await assert.rejects( + discoverLocalAiModels(settings, { fetch: abortAwareFetch(), timeoutMs: 5 }), + /timeout|abort/i, + ); + const controller = new AbortController(); + controller.abort(new DOMException("Cancelled", "AbortError")); + await assert.rejects( + discoverLocalAiModels(settings, { fetch: abortAwareFetch(), signal: controller.signal }), + /cancel|abort/i, + ); +}); + +test("registers startup models with compatibility and refresh replaces them", async () => { + const { root, path } = await configFile(); + try { + let call = 0; + const fetchImpl: typeof fetch = async () => { + call += 1; + return response({ data: [{ id: call === 1 ? "qwen3.8-27b-q4" : "next-chat-model" }] }); + }; + let registered: any; + const pi = { + registerProvider(_id: string, config: unknown) { + registered = config; + }, + on() {}, + }; + await createLocalAiModelsExtension({ configPath: path, fetch: fetchImpl })(pi as any); + assert.deepEqual(registered.models.map((model: any) => model.id), ["qwen3.8-27b-q4"]); + assert.deepEqual(registered.models[0].compat, { + supportsDeveloperRole: false, + supportsReasoningEffort: false, + maxTokensField: "max_tokens", + }); + const refreshed = await registered.refreshModels({ signal: new AbortController().signal }); + assert.deepEqual(refreshed.map((model: any) => model.id), ["next-chat-model"]); + assert.deepEqual(refreshed[0].compat, registered.models[0].compat); + } finally { + await rm(root, { recursive: true, force: true }); + } +}); + +test("Pi startup composes discovered models with configured overrides", async () => { + const root = await mkdtemp(join(tmpdir(), "localai-pi-integration-")); + const server = createServer((request, responseStream) => { + if (request.url === "/v1/models") { + responseStream.writeHead(200, { "content-type": "application/json" }); + responseStream.end(JSON.stringify({ + data: [ + { id: "qwen3.8-27b-heretic-abliterated-uncensored" }, + { id: "qwen3.8-27b-q4" }, + ], + })); + return; + } + responseStream.writeHead(404).end(); + }); + await new Promise((resolve) => server.listen(0, "127.0.0.1", resolve)); + try { + const address = server.address() as { port: number }; + const extensionPath = join(dirname(fileURLToPath(import.meta.url)), "index.ts"); + await writeFile(join(root, "settings.json"), JSON.stringify({ extensions: [extensionPath] })); + await writeFile(join(root, "models.json"), JSON.stringify({ + providers: { + localai: { + baseUrl: `http://127.0.0.1:${address.port}/v1`, + api: "openai-completions", + apiKey: "local-test", + compat: { + supportsDeveloperRole: false, + supportsReasoningEffort: false, + maxTokensField: "max_tokens", + }, + modelOverrides: { + "qwen3.8-27b-heretic-abliterated-uncensored": { + name: "Overridden Heretic Name", + input: ["text", "image"], + }, + }, + }, + }, + })); + const models = (await runPiCatalog(root)).filter((model) => model.provider === "localai"); + assert.deepEqual(models.map((model) => model.id).sort(), [ + "qwen3.8-27b-heretic-abliterated-uncensored", + "qwen3.8-27b-q4", + ]); + const overridden = models.find((model) => model.id.includes("heretic")); + assert.equal(overridden.name, "Overridden Heretic Name"); + assert.deepEqual(overridden.input, ["text", "image"]); + assert.deepEqual(overridden.compat, { + supportsDeveloperRole: false, + supportsReasoningEffort: false, + maxTokensField: "max_tokens", + }); + } finally { + await new Promise((resolve) => server.close(() => resolve())); + await rm(root, { recursive: true, force: true }); + } +}); + +test("successful startup followed by outage returns an empty refreshed catalog", async () => { + const { root, path } = await configFile(); + const originalWarn = console.warn; + console.warn = () => {}; + try { + let call = 0; + let registered: any; + const pi = { + registerProvider(_id: string, config: unknown) { + registered = config; + }, + on() {}, + }; + await createLocalAiModelsExtension({ + configPath: path, + fetch: async () => { + call += 1; + if (call === 1) return response({ data: [{ id: "qwen3.8-27b-q4" }] }); + throw new Error("offline"); + }, + })(pi as any); + assert.deepEqual(registered.models.map((model: any) => model.id), ["qwen3.8-27b-q4"]); + assert.deepEqual(await registered.refreshModels({ signal: new AbortController().signal }), []); + } finally { + console.warn = originalWarn; + await rm(root, { recursive: true, force: true }); + } +}); + +test("continues with an empty catalog and warning when LocalAI is unavailable", async () => { + const { root, path } = await configFile(); + const originalWarn = console.warn; + const warnings: string[] = []; + console.warn = (message?: unknown) => warnings.push(String(message)); + try { + let registered: any; + let sessionStart: ((event: unknown, ctx: any) => void) | undefined; + const pi = { + registerProvider(_id: string, config: unknown) { + registered = config; + }, + on(event: string, handler: (event: unknown, ctx: any) => void) { + if (event === "session_start") sessionStart = handler; + }, + }; + await createLocalAiModelsExtension({ + configPath: path, + fetch: async () => { throw new Error("offline"); }, + })(pi as any); + assert.deepEqual(registered.models, []); + assert.deepEqual(await registered.refreshModels({ signal: new AbortController().signal }), []); + const notices: unknown[] = []; + sessionStart?.({}, { + hasUI: true, + ui: { notify: (...args: unknown[]) => notices.push(args) }, + }); + assert.equal(notices.length, 1); + assert.ok(warnings.length >= 1); + } finally { + console.warn = originalWarn; + await rm(root, { recursive: true, force: true }); + } +}); diff --git a/.pi/agent/extensions/localai-models/index.ts b/.pi/agent/extensions/localai-models/index.ts new file mode 100644 index 0000000..f5834ba --- /dev/null +++ b/.pi/agent/extensions/localai-models/index.ts @@ -0,0 +1,252 @@ +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();