Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 2 additions & 15 deletions src/lib/onboard.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
4 changes: 3 additions & 1 deletion src/lib/onboard/machine/core-flow-phases.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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: [],
Expand Down Expand Up @@ -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"],
},
Expand Down
2 changes: 2 additions & 0 deletions src/lib/onboard/machine/handlers/sandbox-test-fixtures.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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: [],
Expand Down
2 changes: 2 additions & 0 deletions src/lib/onboard/machine/handlers/sandbox.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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,
Expand Down
2 changes: 2 additions & 0 deletions src/lib/onboard/sandbox-create-intent-resolution.ts
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ import {

export type CompleteSandboxCreateIntentInput<Agent, ResourceProfile> = {
sandboxName: string;
inferenceProvider?: string | null;
enabledChannels: readonly string[] | null;
webSearchConfig: WebSearchConfig | null;
agent: Agent;
Expand Down Expand Up @@ -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,
Expand Down
2 changes: 2 additions & 0 deletions src/lib/onboard/sandbox-create-intent-types.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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[];
Expand All @@ -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<string>;
Expand Down
4 changes: 4 additions & 0 deletions src/lib/onboard/sandbox-create-intent.ts
Original file line number Diff line number Diff line change
Expand Up @@ -130,6 +130,7 @@ export function resolveSandboxCreateMessagingProviderRequests(
export function resolveSandboxCreateIntent({
basePolicyPath,
sandboxName,
inferenceProvider,
channels,
enabledChannels,
disabledChannelNames,
Expand Down Expand Up @@ -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,
Expand Down
12 changes: 6 additions & 6 deletions src/lib/onboard/sandbox-create-plan-materialization.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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<string>();
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);
}

Expand Down
119 changes: 119 additions & 0 deletions src/lib/onboard/sandbox-create-plan.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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"]);
});
});
3 changes: 3 additions & 0 deletions src/lib/onboard/sandbox-create-plan.ts
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,7 @@ export type PrepareSandboxCreatePlanInput = {
basePolicyPath: string;
buildCtx: string;
sandboxName: string;
inferenceProvider?: string | null;
channels: MessagingChannel[];
enabledChannels: string[] | null;
disabledChannelNames: ReadonlySet<string>;
Expand Down Expand Up @@ -99,6 +100,7 @@ export function prepareSandboxCreatePlan({
basePolicyPath,
buildCtx,
sandboxName,
inferenceProvider,
channels,
enabledChannels,
disabledChannelNames,
Expand Down Expand Up @@ -131,6 +133,7 @@ export function prepareSandboxCreatePlan({
const intent = resolveSandboxCreateIntent({
basePolicyPath,
sandboxName,
inferenceProvider,
channels,
enabledChannels,
disabledChannelNames,
Expand Down
Loading