fix(FN-766): load pi extension providers for engine

This commit is contained in:
gsxdsm
2026-04-02 22:18:47 -07:00
parent f9dd3983c0
commit c1b6d1c4a9
2 changed files with 293 additions and 6 deletions

View File

@@ -0,0 +1,182 @@
import { beforeEach, describe, expect, it, vi } from "vitest";
const createAgentSessionMock = vi.fn();
const createCodingToolsMock = vi.fn(() => []);
const createReadOnlyToolsMock = vi.fn(() => []);
const createExtensionRuntimeMock = vi.fn();
const discoverAndLoadExtensionsMock = vi.fn().mockResolvedValue({
runtime: { pendingProviderRegistrations: [] },
errors: [],
});
const packageManagerResolveMock = vi.fn().mockResolvedValue({ extensions: [] });
const findMock = vi.fn();
const registerProviderMock = vi.fn();
const refreshMock = vi.fn();
const settingsManagerCreateMock = vi.fn(() => ({ kind: "settings-manager-create" }));
const setFallbackResolverMock = vi.fn();
const reloadMock = vi.fn(async () => {});
vi.mock("@mariozechner/pi-coding-agent", () => ({
AuthStorage: {
create: () => ({
setFallbackResolver: setFallbackResolverMock,
}),
},
createAgentSession: createAgentSessionMock,
createCodingTools: createCodingToolsMock,
createExtensionRuntime: createExtensionRuntimeMock,
createReadOnlyTools: createReadOnlyToolsMock,
DefaultResourceLoader: class {
async reload() {
await reloadMock();
}
},
DefaultPackageManager: class {
async resolve() {
return packageManagerResolveMock();
}
},
discoverAndLoadExtensions: discoverAndLoadExtensionsMock,
getAgentDir: () => "/mock-agent-dir",
ModelRegistry: class {
find(provider: string, modelId: string) {
return findMock(provider, modelId);
}
registerProvider(name: string, config: unknown) {
return registerProviderMock(name, config);
}
refresh() {
return refreshMock();
}
},
SessionManager: {
inMemory: () => ({ kind: "session-manager" }),
},
SettingsManager: {
create: settingsManagerCreateMock,
inMemory: () => ({ kind: "settings-manager" }),
},
}));
describe("createKbAgent", () => {
beforeEach(() => {
vi.clearAllMocks();
findMock.mockImplementation((provider: string, modelId: string) => ({ provider, id: modelId }));
createAgentSessionMock.mockResolvedValue({
session: {
prompt: vi.fn(),
subscribe: vi.fn(),
dispose: vi.fn(),
setThinkingLevel: vi.fn(),
},
});
});
it("registers extension providers before resolving configured models", async () => {
packageManagerResolveMock.mockResolvedValueOnce({
extensions: [{ enabled: true, path: "/extensions/zai-provider" }],
});
discoverAndLoadExtensionsMock.mockResolvedValueOnce({
runtime: {
pendingProviderRegistrations: [
{
name: "zai",
config: { models: [{ id: "glm-5.1" }] },
extensionPath: "/extensions/zai-provider",
},
],
},
errors: [],
});
const { createKbAgent } = await import("./pi.js");
await createKbAgent({
cwd: "/tmp",
systemPrompt: "test",
tools: "readonly",
defaultProvider: "zai",
defaultModelId: "glm-5.1",
});
expect(discoverAndLoadExtensionsMock).toHaveBeenCalledWith(["/extensions/zai-provider"], "/tmp", undefined);
expect(registerProviderMock).toHaveBeenCalledWith("zai", expect.objectContaining({
models: [{ id: "glm-5.1" }],
}));
expect(refreshMock).toHaveBeenCalled();
});
it("avoids lock-based SettingsManager.create when loading extension providers", async () => {
const { createKbAgent } = await import("./pi.js");
await createKbAgent({
cwd: "/tmp",
systemPrompt: "test",
tools: "readonly",
defaultProvider: "openai-codex",
defaultModelId: "gpt-5.4",
});
expect(packageManagerResolveMock).toHaveBeenCalled();
expect(discoverAndLoadExtensionsMock).toHaveBeenCalled();
expect(createAgentSessionMock).toHaveBeenCalledTimes(1);
expect(settingsManagerCreateMock).not.toHaveBeenCalled();
});
it("throws when the configured primary model cannot be resolved", async () => {
findMock.mockImplementation((provider: string, modelId: string) => (
provider === "zai" && modelId === "glm-5.1" ? undefined : { provider, id: modelId }
));
const { createKbAgent } = await import("./pi.js");
await expect(createKbAgent({
cwd: "/tmp",
systemPrompt: "test",
tools: "readonly",
defaultProvider: "zai",
defaultModelId: "glm-5.1",
})).rejects.toThrow("Configured primary model zai/glm-5.1 was not found");
expect(createAgentSessionMock).not.toHaveBeenCalled();
});
it("throws when the configured fallback model cannot be resolved", async () => {
findMock.mockImplementation((provider: string, modelId: string) => (
provider === "openai-codex" && modelId === "missing-model" ? undefined : { provider, id: modelId }
));
const { createKbAgent } = await import("./pi.js");
await expect(createKbAgent({
cwd: "/tmp",
systemPrompt: "test",
tools: "coding",
defaultProvider: "openai-codex",
defaultModelId: "gpt-5.4",
fallbackProvider: "openai-codex",
fallbackModelId: "missing-model",
})).rejects.toThrow("Configured fallback model openai-codex/missing-model was not found");
expect(createAgentSessionMock).not.toHaveBeenCalled();
});
it("creates a session when configured models resolve successfully", async () => {
const { createKbAgent } = await import("./pi.js");
await createKbAgent({
cwd: "/tmp",
systemPrompt: "test",
tools: "readonly",
defaultProvider: "openai-codex",
defaultModelId: "gpt-5.4",
fallbackProvider: "openai-codex",
fallbackModelId: "gpt-5.3-codex",
});
expect(createAgentSessionMock).toHaveBeenCalledTimes(1);
expect(createAgentSessionMock.mock.calls[0][0]).toMatchObject({
model: { provider: "openai-codex", id: "gpt-5.4" },
});
});
});

View File

@@ -5,12 +5,18 @@
* Provides factory functions for creating triage and executor agent sessions.
*/
import { existsSync, readFileSync } from "node:fs";
import { join } from "node:path";
import {
AuthStorage,
createAgentSession,
createCodingTools,
createExtensionRuntime,
createReadOnlyTools,
DefaultResourceLoader,
DefaultPackageManager,
discoverAndLoadExtensions,
getAgentDir,
ModelRegistry,
SessionManager,
SettingsManager,
@@ -73,6 +79,27 @@ export interface AgentOptions {
defaultThinkingLevel?: string;
}
function resolveConfiguredModel(
modelRegistry: ModelRegistry,
kind: "primary" | "fallback",
provider?: string,
modelId?: string,
) {
if (!provider || !modelId) {
return undefined;
}
const model = modelRegistry.find(provider, modelId);
if (model) {
return model;
}
throw new Error(
`Configured ${kind} model ${provider}/${modelId} was not found in the pi model registry. ` +
"Open Settings and choose a model from /api/models, or update your pi model configuration.",
);
}
function isRetryableModelSelectionError(message: string): boolean {
const normalized = message.toLowerCase();
return normalized.includes("rate limit")
@@ -84,6 +111,77 @@ function isRetryableModelSelectionError(message: string): boolean {
|| normalized.includes("temporarily unavailable");
}
interface PackageManagerSettingsView {
getGlobalSettings(): Record<string, any>;
getProjectSettings(): Record<string, any>;
getNpmCommand(): string[] | undefined;
}
function readJsonObject(path: string): Record<string, any> {
if (!existsSync(path)) {
return {};
}
try {
const parsed = JSON.parse(readFileSync(path, "utf-8"));
return parsed && typeof parsed === "object" ? parsed as Record<string, any> : {};
} catch {
return {};
}
}
function createReadOnlyPiSettingsView(cwd: string, agentDir: string): PackageManagerSettingsView {
const globalSettings = readJsonObject(join(agentDir, "settings.json"));
const projectSettings = readJsonObject(join(cwd, ".pi", "settings.json"));
const mergedSettings = { ...globalSettings, ...projectSettings };
return {
getGlobalSettings: () => structuredClone(globalSettings),
getProjectSettings: () => structuredClone(projectSettings),
getNpmCommand: () => Array.isArray(mergedSettings.npmCommand)
? [...mergedSettings.npmCommand]
: undefined,
};
}
async function registerExtensionProviders(cwd: string, modelRegistry: ModelRegistry): Promise<void> {
try {
const agentDir = getAgentDir();
const packageManager = new DefaultPackageManager({
cwd,
agentDir,
settingsManager: createReadOnlyPiSettingsView(cwd, agentDir) as any,
});
const resolvedPaths = await packageManager.resolve();
const packageExtensionPaths = resolvedPaths.extensions
.filter((resource) => resource.enabled)
.map((resource) => resource.path);
const extensionsResult = await discoverAndLoadExtensions(packageExtensionPaths, cwd, undefined);
for (const { path, error } of extensionsResult.errors) {
console.log(`[extensions] Failed to load ${path}: ${error}`);
}
for (const { name, config, extensionPath } of extensionsResult.runtime.pendingProviderRegistrations) {
try {
modelRegistry.registerProvider(name, config);
} catch (error) {
const message = error instanceof Error ? error.message : String(error);
console.log(`[extensions] Failed to register provider from ${extensionPath}: ${message}`);
}
}
extensionsResult.runtime.pendingProviderRegistrations = [];
modelRegistry.refresh();
} catch (error) {
const message = error instanceof Error ? error.message : String(error);
console.log(`[extensions] Failed to discover extensions: ${message}`);
createExtensionRuntime();
modelRegistry.refresh();
}
}
/**
* Create a pi agent session configured for kb.
* Reuses the user's existing pi auth and model configuration.
@@ -91,6 +189,7 @@ function isRetryableModelSelectionError(message: string): boolean {
export async function createKbAgent(options: AgentOptions): Promise<AgentResult> {
const authStorage = AuthStorage.create();
const modelRegistry = new ModelRegistry(authStorage);
await registerExtensionProviders(options.cwd, modelRegistry);
const tools =
options.tools === "readonly"
@@ -103,12 +202,18 @@ export async function createKbAgent(options: AgentOptions): Promise<AgentResult>
});
// Resolve explicit model selection if provider and model ID are specified
const selectedModel = options.defaultProvider && options.defaultModelId
? modelRegistry.find(options.defaultProvider, options.defaultModelId)
: undefined;
const fallbackModel = options.fallbackProvider && options.fallbackModelId
? modelRegistry.find(options.fallbackProvider, options.fallbackModelId)
: undefined;
const selectedModel = resolveConfiguredModel(
modelRegistry,
"primary",
options.defaultProvider,
options.defaultModelId,
);
const fallbackModel = resolveConfiguredModel(
modelRegistry,
"fallback",
options.fallbackProvider,
options.fallbackModelId,
);
const resourceLoader = new DefaultResourceLoader({
cwd: options.cwd,