diff --git a/src/lib/onboard.ts b/src/lib/onboard.ts index 0367dd0a0c4..44cf6c265bb 100644 --- a/src/lib/onboard.ts +++ b/src/lib/onboard.ts @@ -2283,21 +2283,8 @@ async function createSandboxWithBaseImageResolution( const extraProviderPlan = createIntent?.extraProviders ? { extraProviders: createIntent.extraProviders, staleExtraProviders: [] } : planRegisteredExtraProviders(GATEWAY_NAME, { runOpenshell }); - const resolvedCreateIntent = - createIntent?.resolved ?? - (await sandboxCreateIntentResolver.resolve({ - sandboxName, - enabledChannels, - webSearchConfig, - agent, - sandboxGpuConfig: effectiveSandboxGpuConfig, - resourceProfile, - hermesToolGateways, - extraProviders: extraProviderPlan.extraProviders, - staleExtraProviders: extraProviderPlan.staleExtraProviders, - ...(createIntent?.reuseRegisteredCredentials ? { reuseRegisteredCredentials: true } : {}), - ...(createIntent?.policyTier !== undefined ? { policyTier: createIntent.policyTier } : {}), - })); + // biome-ignore format: keep src/lib/onboard.ts net-neutral for growth guardrail. + const resolvedCreateIntent = createIntent?.resolved ?? (await sandboxCreateIntentResolver.resolve({ sandboxName, inferenceProvider: provider, enabledChannels, webSearchConfig, agent, sandboxGpuConfig: effectiveSandboxGpuConfig, resourceProfile, hermesToolGateways, extraProviders: extraProviderPlan.extraProviders, staleExtraProviders: extraProviderPlan.staleExtraProviders, ...(createIntent?.reuseRegisteredCredentials ? { reuseRegisteredCredentials: true } : {}), ...(createIntent?.policyTier !== undefined ? { policyTier: createIntent.policyTier } : {}) })); const messagingCapabilities = await sandboxCreateIntentResolver.rebind( { sandboxName, diff --git a/src/lib/onboard/machine/core-flow-phases.test.ts b/src/lib/onboard/machine/core-flow-phases.test.ts index 52b0d150e4d..ba55a402459 100644 --- a/src/lib/onboard/machine/core-flow-phases.test.ts +++ b/src/lib/onboard/machine/core-flow-phases.test.ts @@ -215,8 +215,9 @@ function createPhases( staleExtraProviders: [], })), resolveSandboxCreateIntent: vi.fn( - async ({ sandboxName, extraProviders, staleExtraProviders }) => ({ + async ({ sandboxName, inferenceProvider, extraProviders, staleExtraProviders }) => ({ sandboxName, + inferenceProvider: inferenceProvider ?? null, activeMessagingChannels: [], messagingProviderRequests: [], reusableMessagingProviders: [], @@ -320,6 +321,7 @@ describe("core onboard flow phases", () => { ); expect(createSandbox.mock.calls[0]?.at(-1)).toMatchObject({ resolved: { + inferenceProvider: "nim", extraProviders: ["current-provider"], staleExtraProviders: ["stale-provider"], }, diff --git a/src/lib/onboard/machine/handlers/sandbox-test-fixtures.ts b/src/lib/onboard/machine/handlers/sandbox-test-fixtures.ts index 7bdbcdbd6b9..7fa35726aa3 100644 --- a/src/lib/onboard/machine/handlers/sandbox-test-fixtures.ts +++ b/src/lib/onboard/machine/handlers/sandbox-test-fixtures.ts @@ -126,10 +126,12 @@ export function createDeps( resolveCreateIntent: vi.fn( async (input: { sandboxName: string; + inferenceProvider?: string | null; extraProviders: readonly string[]; staleExtraProviders: readonly string[]; }) => ({ sandboxName: input.sandboxName, + inferenceProvider: input.inferenceProvider ?? null, activeMessagingChannels: [], messagingProviderRequests: [], reusableMessagingProviders: [], diff --git a/src/lib/onboard/machine/handlers/sandbox.ts b/src/lib/onboard/machine/handlers/sandbox.ts index df33d057093..5169e0726b8 100644 --- a/src/lib/onboard/machine/handlers/sandbox.ts +++ b/src/lib/onboard/machine/handlers/sandbox.ts @@ -218,6 +218,7 @@ export interface SandboxStateOptions< ): import("../../extra-provider-reconciliation").ExtraProviderReconciliationPlan; resolveSandboxCreateIntent(input: { sandboxName: string; + inferenceProvider?: string | null; enabledChannels: readonly string[]; webSearchConfig: WebSearchConfig | null; agent: Agent; @@ -1115,6 +1116,7 @@ class SandboxStateFlow< const reuseRegisteredCredentials = this.resumesSandboxPrompts && this.options.resume; const resolved = await this.deps.resolveSandboxCreateIntent({ sandboxName, + inferenceProvider: this.options.provider, enabledChannels: state.selectedMessagingChannels, webSearchConfig: state.webSearchConfig, agent: this.options.agent, diff --git a/src/lib/onboard/sandbox-create-intent-resolution.ts b/src/lib/onboard/sandbox-create-intent-resolution.ts index 7f284c94a42..3bc70fe0524 100644 --- a/src/lib/onboard/sandbox-create-intent-resolution.ts +++ b/src/lib/onboard/sandbox-create-intent-resolution.ts @@ -20,6 +20,7 @@ import { export type CompleteSandboxCreateIntentInput = { sandboxName: string; + inferenceProvider?: string | null; enabledChannels: readonly string[] | null; webSearchConfig: WebSearchConfig | null; agent: Agent; @@ -116,6 +117,7 @@ export function createSandboxCreateIntentResolver< return resolveSandboxCreateIntent({ basePolicyPath: deps.getAgentPolicyPath(input.agent) || deps.defaultPolicyPath, sandboxName: input.sandboxName, + inferenceProvider: input.inferenceProvider, channels: deps.channels, enabledChannels: filterEnabledChannels(input.enabledChannels, input.agent), disabledChannelNames: messaging.disabledChannelNames, diff --git a/src/lib/onboard/sandbox-create-intent-types.ts b/src/lib/onboard/sandbox-create-intent-types.ts index d87945045a8..353ad41d225 100644 --- a/src/lib/onboard/sandbox-create-intent-types.ts +++ b/src/lib/onboard/sandbox-create-intent-types.ts @@ -40,6 +40,7 @@ export type SandboxCreatePolicyRequest = { */ export type SandboxCreateIntent = { readonly sandboxName: string; + readonly inferenceProvider: string | null; readonly activeMessagingChannels: readonly string[]; readonly messagingProviderRequests: readonly SandboxCreateMessagingProviderRequest[]; readonly reusableMessagingProviders: readonly string[]; @@ -58,6 +59,7 @@ export type SandboxCreateIntent = { export type ResolveSandboxCreateIntentInput = { basePolicyPath: string; sandboxName: string; + inferenceProvider?: string | null; channels: readonly MessagingChannel[]; enabledChannels: string[] | null; disabledChannelNames: ReadonlySet; diff --git a/src/lib/onboard/sandbox-create-intent.ts b/src/lib/onboard/sandbox-create-intent.ts index 5783a59f85d..255dd573d66 100644 --- a/src/lib/onboard/sandbox-create-intent.ts +++ b/src/lib/onboard/sandbox-create-intent.ts @@ -130,6 +130,7 @@ export function resolveSandboxCreateMessagingProviderRequests( export function resolveSandboxCreateIntent({ basePolicyPath, sandboxName, + inferenceProvider, channels, enabledChannels, disabledChannelNames, @@ -168,8 +169,11 @@ export function resolveSandboxCreateIntent({ disabledChannelNames, ); + const normalizedInferenceProvider = inferenceProvider?.trim() || null; + return { sandboxName, + inferenceProvider: normalizedInferenceProvider, activeMessagingChannels, messagingProviderRequests: messagingProviderRequests.map((request) => ({ ...request })), reusableMessagingProviders: enabledReusableMessagingProviders, diff --git a/src/lib/onboard/sandbox-create-plan-materialization.ts b/src/lib/onboard/sandbox-create-plan-materialization.ts index f859695d285..cb85f45e852 100644 --- a/src/lib/onboard/sandbox-create-plan-materialization.ts +++ b/src/lib/onboard/sandbox-create-plan-materialization.ts @@ -159,14 +159,14 @@ export function materializeSandboxCreatePlan({ providerChannels, new Set(intent.disabledChannelNames), ); - for (const provider of messagingProviders) { - createArgs.push("--provider", provider); - } + const createProviders = new Set(); + if (intent.inferenceProvider) createProviders.add(intent.inferenceProvider); + for (const provider of messagingProviders) createProviders.add(provider); if (intent.hermesToolGateways.length > 0) { - createArgs.push("--provider", getHermesToolGatewayProviderName(intent.sandboxName)); + createProviders.add(getHermesToolGatewayProviderName(intent.sandboxName)); } - for (const provider of intent.extraProviders) { - if (messagingProviders.includes(provider)) continue; + for (const provider of intent.extraProviders) createProviders.add(provider); + for (const provider of createProviders) { createArgs.push("--provider", provider); } diff --git a/src/lib/onboard/sandbox-create-plan.test.ts b/src/lib/onboard/sandbox-create-plan.test.ts index b93f08233a2..2f11cce07b7 100644 --- a/src/lib/onboard/sandbox-create-plan.test.ts +++ b/src/lib/onboard/sandbox-create-plan.test.ts @@ -639,3 +639,122 @@ describe("prepareSandboxCreatePlan", () => { expect(providerArgs).toEqual(["sandbox-telegram-bridge", "tavily-search"]); }); }); + +describe("selected inference provider attachment (#7171)", () => { + function resolveWithInferenceProvider(inferenceProvider: string | null) { + return resolveSandboxCreateIntent({ + basePolicyPath: "/repo/policy.yaml", + sandboxName: "sandbox", + inferenceProvider, + channels, + enabledChannels: [], + disabledChannelNames: new Set(), + messagingProviderRequests: [], + primaryMessagingCredentialEnvKeys: [], + reusableMessagingChannels: [], + reusableMessagingProviders: [], + hermesToolGateways: [], + sandboxGpuConfig, + gpuCreateArgs: [], + gpuRoutePlan: "native-only", + sandboxGpuLogMessage: null, + policyTier: null, + }); + } + + function planWithInferenceProvider(overrides: { + inferenceProvider?: string | null; + messagingTokenDefs?: MessagingTokenDef[]; + reusableMessagingProviders?: string[]; + extraProviders?: string[]; + hermesToolGateways?: string[]; + upsertMessagingProviders?: () => string[]; + }) { + return prepareSandboxCreatePlan({ + basePolicyPath: "/repo/policy.yaml", + buildCtx: "/tmp/nemoclaw-build-1", + sandboxName: "sandbox", + inferenceProvider: overrides.inferenceProvider, + channels, + enabledChannels: [], + disabledChannelNames: new Set(), + messagingTokenDefs: overrides.messagingTokenDefs ?? [], + reusableMessagingChannels: [], + reusableMessagingProviders: overrides.reusableMessagingProviders ?? [], + extraProviders: overrides.extraProviders ?? [], + hermesToolGateways: overrides.hermesToolGateways ?? [], + sandboxGpuConfig, + gpuRoutePlan: "native-only", + sandboxGpuLogMessage: null, + appendResourceFlags: vi.fn(), + runProviderPreDeleteCleanup: vi.fn(), + upsertMessagingProviders: vi.fn(overrides.upsertMessagingProviders ?? (() => [])), + getMessagingChannelForEnvKey: (envKey) => + envKey === "TELEGRAM_BOT_TOKEN" ? "telegram" : null, + getHermesToolGatewayProviderName: (sandboxName) => `${sandboxName}-hermes-tools`, + deps: { + prepareInitialSandboxCreatePolicy: vi.fn(() => ({ + policyPath: "/tmp/policy.yaml", + appliedPresets: [], + })), + buildSandboxGpuCreateArgs: vi.fn(() => []), + }, + }); + } + + function providerArgsOf(createArgs: readonly string[]): string[] { + return createArgs + .map((arg, index) => (arg === "--provider" ? createArgs[index + 1] : null)) + .filter((value): value is string => value !== null); + } + + it("serializes the selected provider into the intent without a credential value", () => { + const intent = resolveWithInferenceProvider(" nvidia-router "); + expect(intent.inferenceProvider).toBe("nvidia-router"); + expect(JSON.parse(JSON.stringify(intent)).inferenceProvider).toBe("nvidia-router"); + }); + + it("treats a blank selected provider as absent", () => { + expect(resolveWithInferenceProvider(" ").inferenceProvider).toBeNull(); + expect(resolveWithInferenceProvider(null).inferenceProvider).toBeNull(); + }); + + it.each([ + "nvidia-router", + "openai-compatible", + "vllm-local", + ])("attaches the selected provider %s first on create", (provider) => { + const result = planWithInferenceProvider({ + inferenceProvider: provider, + extraProviders: ["tavily-search"], + }); + expect(providerArgsOf(result.createArgs)).toEqual([provider, "tavily-search"]); + }); + + it("emits the selected provider exactly once when it also appears as an extra provider", () => { + const result = planWithInferenceProvider({ + inferenceProvider: "vllm-local", + extraProviders: ["vllm-local", "tavily-search"], + }); + expect(providerArgsOf(result.createArgs)).toEqual(["vllm-local", "tavily-search"]); + }); + + it("emits the selected provider exactly once when it also backs a messaging channel", () => { + const result = planWithInferenceProvider({ + inferenceProvider: "sandbox-telegram-bridge", + messagingTokenDefs: [ + { name: "sandbox-telegram-bridge", envKey: "TELEGRAM_BOT_TOKEN", token: "telegram" }, + ], + upsertMessagingProviders: () => ["sandbox-telegram-bridge"], + }); + expect(providerArgsOf(result.createArgs)).toEqual(["sandbox-telegram-bridge"]); + }); + + it("omits an inference --provider when no provider is selected", () => { + const result = planWithInferenceProvider({ + inferenceProvider: null, + extraProviders: ["tavily-search"], + }); + expect(providerArgsOf(result.createArgs)).toEqual(["tavily-search"]); + }); +}); diff --git a/src/lib/onboard/sandbox-create-plan.ts b/src/lib/onboard/sandbox-create-plan.ts index 9343bc78d2d..b2e7a88d704 100644 --- a/src/lib/onboard/sandbox-create-plan.ts +++ b/src/lib/onboard/sandbox-create-plan.ts @@ -71,6 +71,7 @@ export type PrepareSandboxCreatePlanInput = { basePolicyPath: string; buildCtx: string; sandboxName: string; + inferenceProvider?: string | null; channels: MessagingChannel[]; enabledChannels: string[] | null; disabledChannelNames: ReadonlySet; @@ -99,6 +100,7 @@ export function prepareSandboxCreatePlan({ basePolicyPath, buildCtx, sandboxName, + inferenceProvider, channels, enabledChannels, disabledChannelNames, @@ -131,6 +133,7 @@ export function prepareSandboxCreatePlan({ const intent = resolveSandboxCreateIntent({ basePolicyPath, sandboxName, + inferenceProvider, channels, enabledChannels, disabledChannelNames,