fix(coding-agent): refresh model catalogs in picker

This commit is contained in:
Mario Zechner
2026-07-15 12:35:29 +02:00
parent ff28097a36
commit fab309e955
7 changed files with 126 additions and 8 deletions
+1
View File
@@ -34,6 +34,7 @@
### Fixed
- Fixed configured-provider catalog refresh to parse pi.dev's model-ID keyed responses, treat unimplemented routes as unavailable overlays, and show concise refresh status in `/model`.
- Fixed inherited OpenRouter model context windows to use the top provider's actual context length ([#6481](https://github.com/earendil-works/pi-mono/pull/6481) by [@davidbrai](https://github.com/davidbrai)).
- Fixed inherited OpenRouter OpenAI-compatible session IDs to use the `x-session-id` header instead of OpenAI-specific session-affinity fields ([#6366](https://github.com/earendil-works/pi/issues/6366)).
- Fixed `Ctrl+V` to paste clipboard text when the pasteboard does not contain an image.
@@ -501,10 +501,14 @@ export class ModelRuntime implements Models {
}
async refresh(options: ModelsRefreshOptions = {}): Promise<ModelsRefreshResult> {
const refreshOptions = {
...options,
allowNetwork: options.allowNetwork ?? this.allowModelNetwork,
};
// Published pi-ai builds before ModelsStore returned void and accepted a provider ID.
// The fallback keeps source-mode CLI tests working without rebuilding workspace dependencies.
const result = ((await this.models.refresh(options)) as ModelsRefreshResult | undefined) ?? {
aborted: options.signal?.aborted ?? false,
const result = ((await this.models.refresh(refreshOptions)) as ModelsRefreshResult | undefined) ?? {
aborted: refreshOptions.signal?.aborted ?? false,
errors: new Map(),
};
this.updateModelSnapshot();
@@ -17,6 +17,8 @@ function parseCatalog(providerId: string, value: unknown): Model<Api>[] {
? value
: typeof value === "object" && value !== null && "models" in value && Array.isArray(value.models)
? value.models
: typeof value === "object" && value !== null
? Object.values(value)
: undefined;
if (!entries) throw new Error(`Invalid model catalog for provider "${providerId}"`);
return entries
@@ -44,6 +46,7 @@ export function withRemoteCatalog(provider: Provider, catalogBaseUrl: string = D
headers: { accept: "application/json" },
signal: context.signal,
});
if (response.status === 404 || response.status === 501) return;
if (!response.ok) {
throw new Error(`Model catalog request failed for ${provider.id}: ${response.status}`);
}
@@ -56,6 +56,8 @@ export class ModelSelectorComponent extends Container implements Focusable {
private onSelectCallback: (model: Model<any>) => void;
private onCancelCallback: () => void;
private errorMessage?: string;
private refreshStatusMessage = "Refreshing model catalogs…";
private refreshStatusSuccess = false;
private tui: TUI;
private scopedModels: ReadonlyArray<ScopedModelItem>;
private scope: ModelScope = "all";
@@ -167,12 +169,19 @@ export class ModelSelectorComponent extends Container implements Focusable {
try {
const result = await this.modelRuntime.refresh({ signal: this.refreshAbortController.signal });
if (this.closed) return;
this.refreshStatusMessage = "";
if (result.aborted && timedOut) {
this.errorMessage = "Model refresh timed out; showing cached models.";
} else if (result.errors.size > 0) {
this.errorMessage = `Model refresh failed for: ${[...result.errors.keys()].join(", ")}`;
} else if (result.errors.size === 1) {
this.errorMessage = `Could not refresh ${result.errors.keys().next().value}; showing cached models.`;
} else if (result.errors.size > 1) {
this.errorMessage = `Could not refresh ${result.errors.size} model catalogs; showing cached models.`;
} else {
this.errorMessage = this.modelRuntime.getError();
if (!this.errorMessage) {
this.refreshStatusMessage = "Model catalogs refreshed.";
this.refreshStatusSuccess = true;
}
}
this.loadModelsFromSnapshot();
this.filterModels(this.searchInput.getValue());
@@ -288,6 +297,12 @@ export class ModelSelectorComponent extends Container implements Focusable {
this.listContainer.addChild(new Spacer(1));
this.listContainer.addChild(new Text(theme.fg("muted", ` Model Name: ${selected.model.name}`), 0, 0));
}
if (this.refreshStatusMessage) {
this.listContainer.addChild(new Spacer(1));
this.listContainer.addChild(
new Text(theme.fg(this.refreshStatusSuccess ? "success" : "muted", ` ${this.refreshStatusMessage}`), 0, 0),
);
}
}
handleInput(keyData: string): void {
@@ -11,11 +11,11 @@ function wrap(runtime: ModelRuntime): ModelRegistry {
}
export async function createModelRegistry(credentials: CredentialStore, modelsPath?: string): Promise<ModelRegistry> {
return wrap(await ModelRuntime.create({ credentials, modelsPath }));
return wrap(await ModelRuntime.create({ credentials, modelsPath, allowModelNetwork: false }));
}
export async function createInMemoryModelRegistry(credentials: CredentialStore): Promise<ModelRegistry> {
return wrap(await ModelRuntime.create({ credentials, modelsPath: null }));
return wrap(await ModelRuntime.create({ credentials, modelsPath: null, allowModelNetwork: false }));
}
export function getModelRuntime(modelRegistry: ModelRegistry): ModelRuntime {
@@ -0,0 +1,93 @@
import { createProvider, InMemoryModelsStore, type Model } from "@earendil-works/pi-ai";
import { afterEach, describe, expect, it, vi } from "vitest";
import { withRemoteCatalog } from "../src/core/remote-catalog-provider.ts";
function model(id: string): Model<"openai-completions"> {
return {
id,
name: id,
api: "openai-completions",
provider: "test-provider",
baseUrl: "https://example.test/v1",
reasoning: false,
input: ["text"],
cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 },
contextWindow: 1000,
maxTokens: 100,
};
}
afterEach(() => vi.restoreAllMocks());
describe("remote catalog provider", () => {
it("parses catalogs keyed by model ID", async () => {
vi.spyOn(globalThis, "fetch").mockResolvedValue(
new Response(JSON.stringify({ dynamic: model("dynamic") }), {
status: 200,
headers: { "content-type": "application/json" },
}),
);
const provider = withRemoteCatalog(
createProvider({
id: "test-provider",
auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } },
models: [model("static")],
api: {
stream: () => {
throw new Error("not used");
},
streamSimple: () => {
throw new Error("not used");
},
},
}),
);
const store = new InMemoryModelsStore();
await provider.refreshModels?.({
credential: { type: "api_key" },
store: {
read: () => store.read(provider.id),
write: (models) => store.write(provider.id, models),
delete: () => store.delete(provider.id),
},
allowNetwork: true,
});
expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "dynamic"]);
expect((await store.read(provider.id))?.map((entry) => entry.id)).toEqual(["dynamic"]);
});
it("treats unimplemented pi.dev catalog routes as an unavailable overlay", async () => {
vi.spyOn(globalThis, "fetch").mockResolvedValue(new Response("not implemented", { status: 501 }));
const provider = withRemoteCatalog(
createProvider({
id: "test-provider",
auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } },
models: [model("static")],
api: {
stream: () => {
throw new Error("not used");
},
streamSimple: () => {
throw new Error("not used");
},
},
}),
);
const store = new InMemoryModelsStore();
await expect(
provider.refreshModels?.({
credential: { type: "api_key" },
store: {
read: () => store.read(provider.id),
write: (models) => store.write(provider.id, models),
delete: () => store.delete(provider.id),
},
allowNetwork: true,
}),
).resolves.toBeUndefined();
expect(provider.getModels().map((entry) => entry.id)).toEqual(["static"]);
expect(await store.read(provider.id)).toBeUndefined();
});
});
@@ -86,7 +86,9 @@ describe("issue #3217 scoped model ordering", () => {
);
await vi.waitFor(() => {
expect(stripAnsi(selector.render(120).join("\n"))).toContain(`[${modelOne.provider}]`);
const rendered = stripAnsi(selector.render(120).join("\n"));
expect(rendered).toContain(`[${modelOne.provider}]`);
expect(rendered).toContain("Model catalogs refreshed.");
});
const renderedLines = stripAnsi(selector.render(120).join("\n"))