Files
dotfiles/.pi/agent/extensions/localai-models/index.test.ts
T

368 lines
14 KiB
TypeScript

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<Response>((_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<any[]>((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<HeadersInit | undefined> = [];
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<void>((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<void>((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 });
}
});