fix(coding-agent): select Radius models after catalog discovery

This commit is contained in:
Armin Ronacher
2026-09-06 00:30:17 +02:00
parent fcff255b00
commit 9767ba275f
4 changed files with 167 additions and 47 deletions
+4
View File
@@ -6,6 +6,10 @@
- Enabled strict-prefer JSON-schema sampling by default for built-in `read`, `bash`, `powershell`, `edit`, and `write` tools, without requiring `PI_EXPERIMENTAL`. Extensions can re-register tool definitions with `constrainedSampling: false`.
### Fixed
- Fixed premature missing-model errors after login by waiting for catalog discovery. Radius now defaults to `balanced`, falling back to the first available Radius model when needed.
## [0.85.1] - 2026-09-05
### New Features
@@ -24,7 +24,7 @@ export const defaultModelPerProvider: Record<KnownProvider, string> = {
openai: "gpt-5.5",
"azure-openai-responses": "gpt-5.4",
"openai-codex": "gpt-5.5",
radius: "auto",
radius: "balanced",
nvidia: "nvidia/nemotron-3-super-120b-a12b",
deepseek: "deepseek-v4-pro",
google: "gemini-3.1-pro-preview",
@@ -5678,61 +5678,83 @@ export class InteractiveMode {
): Promise<void> {
const actionLabel = authType === "oauth" ? `Logged in to ${providerName}` : `Saved API key for ${providerName}`;
let selectedModel: Model<any> | undefined;
let selectionError: string | undefined;
if (isUnknownModel(previousModel)) {
const availableModels = this.session.modelRuntime.getAvailableSnapshot();
const providerModels = availableModels.filter((model) => model.provider === providerId);
// Matches LLAMA_PROVIDER_ID from extensions/llama/provider.ts; kept inline to avoid coupling interactive mode to the built-in extension.
if (providerId === "llama.cpp") {
selectionError = llamaCppPostLoginGuidance(actionLabel, providerModels.length);
} else if (!hasDefaultModelProvider(providerId)) {
selectionError = `${actionLabel}, but no default model is configured for provider "${providerId}". Use /model to select a model.`;
} else if (providerModels.length === 0) {
selectionError = `${actionLabel}, but no models are available for that provider. Use /model to select a model.`;
} else {
const defaultModelId = defaultModelPerProvider[providerId];
selectedModel = providerModels.find((model) => model.id === defaultModelId);
if (!selectedModel) {
selectionError = `${actionLabel}, but its default model "${defaultModelId}" is not available. Use /model to select a model.`;
const session = this.session;
// Dynamic catalogs may be empty until the first authenticated network refresh.
const deferSelection =
isUnknownModel(previousModel) &&
hasDefaultModelProvider(providerId) &&
!session.modelRuntime
.getAvailableSnapshot()
.some((model) => model.provider === providerId && model.id === defaultModelPerProvider[providerId]);
const finishAuthentication = async () => {
let selectedModel: Model<any> | undefined;
let selectionError: string | undefined;
if (isUnknownModel(previousModel)) {
const availableModels = this.session.modelRuntime.getAvailableSnapshot();
const providerModels = availableModels.filter((model) => model.provider === providerId);
// Matches LLAMA_PROVIDER_ID from extensions/llama/provider.ts; kept inline to avoid coupling interactive mode to the built-in extension.
if (providerId === "llama.cpp") {
selectionError = llamaCppPostLoginGuidance(actionLabel, providerModels.length);
} else if (!hasDefaultModelProvider(providerId)) {
selectionError = `${actionLabel}, but no default model is configured for provider "${providerId}". Use /model to select a model.`;
} else if (providerModels.length === 0) {
selectionError = `${actionLabel}, but no models are available for that provider. Use /model to select a model.`;
} else {
try {
await this.session.setModel(selectedModel, { persist: true });
} catch (error: unknown) {
selectedModel = undefined;
const errorMessage = error instanceof Error ? error.message : String(error);
selectionError = `${actionLabel}, but selecting its default model failed: ${errorMessage}. Use /model to select a model.`;
const defaultModelId = defaultModelPerProvider[providerId];
// Radius catalogs vary by account; prefer balanced, then use catalog order.
selectedModel =
providerModels.find((model) => model.id === defaultModelId) ??
(providerId === "radius" ? providerModels[0] : undefined);
if (!selectedModel) {
selectionError = `${actionLabel}, but its default model "${defaultModelId}" is not available. Use /model to select a model.`;
} else {
try {
await this.session.setModel(selectedModel, { persist: true });
} catch (error: unknown) {
selectedModel = undefined;
const errorMessage = error instanceof Error ? error.message : String(error);
selectionError = `${actionLabel}, but selecting its default model failed: ${errorMessage}. Use /model to select a model.`;
}
}
}
}
}
await this.updateAvailableProviderCount();
this.footer.invalidate();
this.updateEditorBorderColor();
if (selectedModel) {
this.showStatus(`${actionLabel}. Selected ${selectedModel.id}. Credentials saved to ${getAuthPath()}`);
void this.maybeWarnAboutAnthropicSubscriptionAuth(selectedModel);
this.checkDaxnutsEasterEgg(selectedModel);
} else {
this.showStatus(`${actionLabel}. Credentials saved to ${getAuthPath()}`);
if (selectionError) {
this.showError(selectionError);
await this.updateAvailableProviderCount();
this.footer.invalidate();
this.updateEditorBorderColor();
if (selectedModel) {
this.showStatus(`${actionLabel}. Selected ${selectedModel.id}. Credentials saved to ${getAuthPath()}`);
void this.maybeWarnAboutAnthropicSubscriptionAuth(selectedModel);
this.checkDaxnutsEasterEgg(selectedModel);
} else {
void this.maybeWarnAboutAnthropicSubscriptionAuth();
this.showStatus(`${actionLabel}. Credentials saved to ${getAuthPath()}`);
if (selectionError) {
this.showError(selectionError);
} else {
void this.maybeWarnAboutAnthropicSubscriptionAuth();
}
}
};
if (deferSelection) {
this.showStatus(`${actionLabel}. Credentials saved to ${getAuthPath()}. Refreshing model catalog…`);
} else {
await finishAuthentication();
}
const controller = new AbortController();
const timeout = setTimeout(() => controller.abort(), 15_000);
void this.session.modelRuntime
void session.modelRuntime
.refresh({ providers: [providerId], signal: controller.signal })
.then((result) => {
.then(async (result) => {
if (result.aborted) {
this.showWarning(`${actionLabel}, but its model catalog refresh timed out; using cached models.`);
} else if (result.errors.size > 0) {
this.showWarning(`${actionLabel}, but its model catalog could not be refreshed; using cached models.`);
}
// Do not replace a model or session selected while the refresh was running.
if (deferSelection && this.session === session && session.model === previousModel) {
await finishAuthentication();
}
this.updateAvailableProviderCount();
this.footer.invalidate();
this.ui.requestRender();
@@ -1,10 +1,19 @@
import type { Api, Model, Provider } from "@earendil-works/pi-ai";
import { afterEach, describe, expect, it, vi } from "vitest";
import { AuthStorage } from "../../../src/core/auth-storage.ts";
import { defaultModelPerProvider } from "../../../src/core/model-resolver.ts";
import { ModelRuntime } from "../../../src/core/model-runtime.ts";
import { InteractiveMode } from "../../../src/modes/interactive/interactive-mode.ts";
import { createHarness, type Harness } from "../harness.ts";
const complete = Reflect.get(InteractiveMode.prototype, "completeProviderAuthentication") as (
this: object,
providerId: string,
providerName: string,
authType: "oauth" | "api_key",
previousModel: Model<Api>,
) => Promise<void>;
const dynamicModel: Model<"openai-completions"> = {
id: "dynamic",
name: "Dynamic",
@@ -102,14 +111,6 @@ describe("issues #7027 and #7113 credential refresh hang", () => {
checkDaxnutsEasterEgg: vi.fn(),
ui: { requestRender: vi.fn() },
};
const complete = Reflect.get(InteractiveMode.prototype, "completeProviderAuthentication") as (
this: object,
providerId: string,
providerName: string,
authType: "oauth" | "api_key",
previousModel: Model<Api>,
) => Promise<void>;
await complete.call(context, dynamicModel.provider, "Stalled Login", "api_key", harness.getModel());
expect(runtime.refresh).toHaveBeenCalledWith({
providers: [dynamicModel.provider],
@@ -123,3 +124,96 @@ describe("issues #7027 and #7113 credential refresh hang", () => {
);
});
});
describe("post-login model discovery", () => {
let harness: Harness | undefined;
afterEach(() => {
vi.useRealTimers();
harness?.cleanup();
vi.restoreAllMocks();
});
async function startLogin() {
harness = await createHarness();
vi.useFakeTimers();
const session = harness.session;
const model = harness.getModel();
const unknownModel = { ...model, id: "unknown", provider: "unknown", api: "unknown" };
const currentModel = vi.spyOn(session, "model", "get").mockReturnValue(unknownModel);
const availableModels = vi.spyOn(session.modelRuntime, "getAvailableSnapshot").mockReturnValue([]);
let finishRefresh = () => {};
vi.spyOn(session.modelRuntime, "refresh").mockImplementation(
(options) =>
new Promise((resolve) => {
finishRefresh = () => resolve({ aborted: false, errors: new Map() });
options?.signal?.addEventListener("abort", () => resolve({ aborted: true, errors: new Map() }), {
once: true,
});
}),
);
const setModel = vi.spyOn(session, "setModel").mockResolvedValue();
const context = {
session,
updateAvailableProviderCount: vi.fn(),
footer: { invalidate: vi.fn() },
updateEditorBorderColor: vi.fn(),
showStatus: vi.fn(),
showError: vi.fn(),
showWarning: vi.fn(),
maybeWarnAboutAnthropicSubscriptionAuth: vi.fn(),
checkDaxnutsEasterEgg: vi.fn(),
ui: { requestRender: vi.fn() },
};
await complete.call(context, "radius", "Radius", "oauth", unknownModel);
expect(context.showStatus).toHaveBeenCalledWith(expect.stringContaining("Credentials saved"));
expect(context.showError).not.toHaveBeenCalled();
expect(setModel).not.toHaveBeenCalled();
return {
...context,
setModel,
currentModel,
async discover(ids: string[]) {
availableModels.mockReturnValue(ids.map((id) => ({ ...model, provider: "radius", id })));
finishRefresh();
await vi.advanceTimersByTimeAsync(0);
},
};
}
it.each([
{ models: ["fast", "balanced"], selected: "balanced" },
{ models: ["fast", "powerful"], selected: "fast" },
])("selects $selected from the refreshed catalog $models", async ({ models, selected }) => {
expect(defaultModelPerProvider.radius).toBe("balanced");
const login = await startLogin();
await login.discover(models);
expect(login.setModel).toHaveBeenCalledWith(expect.objectContaining({ provider: "radius", id: selected }), {
persist: true,
});
expect(login.showError).not.toHaveBeenCalled();
});
it("reports an empty catalog only after refresh", async () => {
const login = await startLogin();
await login.discover([]);
expect(login.setModel).not.toHaveBeenCalled();
expect(login.showError).toHaveBeenCalledWith(expect.stringContaining("no models are available"));
});
it("preserves a model selected during refresh", async () => {
const login = await startLogin();
login.currentModel.mockReturnValue(harness!.getModel());
await login.discover(["fast", "balanced"]);
expect(login.setModel).not.toHaveBeenCalled();
expect(login.showError).not.toHaveBeenCalled();
});
it("bounds refresh to 15 seconds", async () => {
const login = await startLogin();
await vi.advanceTimersByTimeAsync(15_000);
expect(login.showWarning).toHaveBeenCalledWith(expect.stringContaining("timed out"));
expect(login.showError).toHaveBeenCalledWith(expect.stringContaining("no models are available"));
expect(login.setModel).not.toHaveBeenCalled();
});
});