Skip to content
11 changes: 4 additions & 7 deletions src/lib/onboard.ts
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@ const {
createNvidiaFeaturedModelSession,
createRemoteModelValidator,
resolveCompatibleEndpointSelection,
selectFeaturedModelAfterCredentialPrompt,
}: typeof import("./onboard/setup-nim-selection") = require("./onboard/setup-nim-selection");
const setupNimFlow: typeof import("./onboard/setup-nim-flow") = require("./onboard/setup-nim-flow");
const openrouterSelection: typeof import("./onboard/openrouter-selection") = require("./onboard/openrouter-selection");
Expand Down Expand Up @@ -2539,6 +2540,7 @@ async function handleRemoteProviderSelection(args: RemoteProviderSelectionArgs,
hydrateCredentialEnv(state.credentialEnv);
if (selected.key === "build") {
providerKeyBridge.stageBuildProviderKeyBridge();
let apiKeyNavigation: unknown = null;
if (isNonInteractive()) {
const reuseGatewayCredential = buildCredentialReuse.resolveNonInteractiveBuildCredential({
provider: state.provider,
Expand All @@ -2549,14 +2551,9 @@ async function handleRemoteProviderSelection(args: RemoteProviderSelectionArgs,
state.skipHostInferenceSmoke = reuseGatewayCredential;
state.reuseGatewayCredentialWithoutLocalKey = reuseGatewayCredential;
} else {
await ensureApiKey();
apiKeyNavigation = await ensureApiKey();
}
state.model = await state.nvidiaFeaturedModels!.select(
requestedModel || (typeof state.model === "string" ? state.model : null),
recoveredFromSandbox ? recoveredModel : null,
isNonInteractive(),
process.env.NEMOCLAW_MODEL,
);
state.model = await selectFeaturedModelAfterCredentialPrompt(state.nvidiaFeaturedModels!, apiKeyNavigation, credentialPrompt.shouldReturnToProviderSelection, requestedModel || (typeof state.model === "string" ? state.model : null), recoveredFromSandbox ? recoveredModel : null, isNonInteractive(), process.env.NEMOCLAW_MODEL);
if (isBackToSelection(state.model)) {
console.log(" Returning to provider selection.");
console.log("");
Expand Down
54 changes: 53 additions & 1 deletion src/lib/onboard/nvidia-featured-model-selection.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,12 @@ import { beforeEach, describe, expect, it, vi } from "vitest";

import { promptCloudModel } from "../inference/model-prompts";
import { BACK_TO_SELECTION } from "../navigation";
import { createNvidiaFeaturedModelSession } from "./nvidia-featured-model-selection";
import { shouldReturnToProviderSelection } from "./credential-navigation";
import {
createNvidiaFeaturedModelSession,
type NvidiaFeaturedModelSession,
selectFeaturedModelAfterCredentialPrompt,
} from "./nvidia-featured-model-selection";

vi.mock("../inference/model-prompts", () => ({
promptCloudModel: vi.fn(),
Expand Down Expand Up @@ -99,4 +104,51 @@ describe("NVIDIA featured model selection", () => {
"environment/model",
);
});

it("skips the catalog when the NVIDIA API key prompt asks to go back (#9404)", async () => {
const select = vi.fn().mockResolvedValue("nvidia/selected-model");
const session = { select } as unknown as NvidiaFeaturedModelSession;
const exitOnboard = vi.fn(() => {
throw new Error("exit onboarding");
}) as unknown as () => never;
const shouldReturn = (result: unknown) => shouldReturnToProviderSelection(result, exitOnboard);

await expect(
selectFeaturedModelAfterCredentialPrompt(
session,
{ kind: "back" },
shouldReturn,
null,
null,
false,
),
).resolves.toBe(BACK_TO_SELECTION);
expect(select).not.toHaveBeenCalled();
expect(exitOnboard).not.toHaveBeenCalled();

await expect(
selectFeaturedModelAfterCredentialPrompt(
session,
{ kind: "credential", value: "nvapi-good" },
shouldReturn,
null,
null,
true,
"env/model",
),
).resolves.toBe("nvidia/selected-model");
expect(select).toHaveBeenCalledWith(null, null, true, "env/model");

await expect(
selectFeaturedModelAfterCredentialPrompt(
session,
{ kind: "exit" },
shouldReturn,
null,
null,
false,
),
).rejects.toThrow("exit onboarding");
expect(select).toHaveBeenCalledTimes(1);
});
});
18 changes: 18 additions & 0 deletions src/lib/onboard/nvidia-featured-model-selection.ts
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import {
createNvidiaFeaturedModelPromptOptionsLoader,
type NvidiaFeaturedModelOptions,
} from "../inference/nvidia-featured-models";
import { BACK_TO_SELECTION } from "../navigation";

export type NvidiaFeaturedModelSession = {
select: (
Expand Down Expand Up @@ -57,3 +58,20 @@ export function createNvidiaFeaturedModelSession(
},
};
}

/**
* Select a featured model only when the credential prompt did not ask to leave. `back` at the
* API key prompt must return to provider selection instead of loading the catalog (#9404).
*/
export async function selectFeaturedModelAfterCredentialPrompt(
session: NvidiaFeaturedModelSession,
credentialNavigation: unknown,
shouldReturnToProviderSelection: (result: unknown) => boolean,
requestedModel: string | null,
recoveredModel: string | null,
nonInteractive: boolean,
envModel?: string,
): Promise<ModelPromptResult> {
if (shouldReturnToProviderSelection(credentialNavigation)) return BACK_TO_SELECTION;
return session.select(requestedModel, recoveredModel, nonInteractive, envModel);
}
5 changes: 4 additions & 1 deletion src/lib/onboard/setup-nim-selection.ts
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,10 @@ import type { NvidiaFeaturedModelSession } from "./nvidia-featured-model-selecti
import { exitOnboardFromPrompt, getNavigationChoice } from "./prompt-helpers";
import type { ReasoningEffort } from "./reasoning-mode";

export { createNvidiaFeaturedModelSession } from "./nvidia-featured-model-selection";
export {
createNvidiaFeaturedModelSession,
selectFeaturedModelAfterCredentialPrompt,
} from "./nvidia-featured-model-selection";

export type SetupNimSelectionBackNavigation = Readonly<{ kind: "NEMOCLAW_BACK_TO_SELECTION" }>;

Expand Down
Loading