feat(pi): discover LocalAI models dynamically
This commit is contained in:
@@ -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<string, unknown>;
|
||||
};
|
||||
type ModelsFile = {
|
||||
providers?: Record<string, {
|
||||
baseUrl?: unknown;
|
||||
apiKey?: unknown;
|
||||
compat?: unknown;
|
||||
}>;
|
||||
};
|
||||
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<string, unknown>;
|
||||
};
|
||||
|
||||
function agentDirectory() {
|
||||
return process.env.PI_CODING_AGENT_DIR ?? join(homedir(), ".pi", "agent");
|
||||
}
|
||||
|
||||
function runShellCommand(command: string, timeoutMs: number, signal?: AbortSignal) {
|
||||
return new Promise<string>((resolve, reject) => {
|
||||
const child = spawn("/bin/sh", ["-lc", command], {
|
||||
detached: true,
|
||||
stdio: ["ignore", "pipe", "pipe"],
|
||||
});
|
||||
let stdout = "";
|
||||
let settled = false;
|
||||
let forceKill: ReturnType<typeof setTimeout> | 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<string, unknown>,
|
||||
): LocalAiModel[] {
|
||||
const seen = new Set<string>();
|
||||
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<string, unknown>) }
|
||||
: 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();
|
||||
Reference in New Issue
Block a user