diff --git a/src/lib/onboard/machine/flow-context.test.ts b/src/lib/onboard/machine/flow-context.test.ts index 8ddd580d298..c2cf9d51dd1 100644 --- a/src/lib/onboard/machine/flow-context.test.ts +++ b/src/lib/onboard/machine/flow-context.test.ts @@ -5,6 +5,7 @@ import { describe, expect, it } from "vitest"; import { createSession } from "../../state/onboard-session"; import { + assertProviderModelSelectedContext, assertProviderSelectedContext, assertSandboxCreatedContext, mergeOnboardFlowContext, @@ -89,6 +90,17 @@ describe("onboard flow context helpers", () => { }); }); + it("asserts provider/model-selected context before consumers use provider output", () => { + const context = mergeOnboardFlowContext(baseContext(), { + provider: "nvidia-prod", + model: "model", + }); + + expect(() => + assertProviderModelSelectedContext(context, "provider inference result"), + ).not.toThrow(); + }); + it("asserts provider-selected context before sandbox setup", () => { const context = mergeOnboardFlowContext(baseContext(), { provider: "nvidia-prod", @@ -98,6 +110,12 @@ describe("onboard flow context helpers", () => { expect(() => assertProviderSelectedContext(context, "sandbox setup")).not.toThrow(); }); + it("rejects missing provider/model-selected context fields", () => { + expect(() => + assertProviderModelSelectedContext(baseContext(), "provider inference result"), + ).toThrow(/Onboarding state is incomplete before provider inference result\./); + }); + it("rejects missing provider-selected context fields", () => { expect(() => assertProviderSelectedContext(baseContext(), "sandbox setup")).toThrow( /Onboarding state is incomplete before sandbox setup\./, diff --git a/src/lib/onboard/machine/flow-context.ts b/src/lib/onboard/machine/flow-context.ts index bb7ace415ab..4fb20a38fdf 100644 --- a/src/lib/onboard/machine/flow-context.ts +++ b/src/lib/onboard/machine/flow-context.ts @@ -77,11 +77,21 @@ export interface SandboxCreatedContextUpdate { webSearchSupported: boolean; } +export function assertProviderModelSelectedContext( + context: Context, + stepName: string, +): asserts context is ProviderModelSelectedOnboardFlowContext { + if (!context.model || !context.provider) { + throw new Error(`Onboarding state is incomplete before ${stepName}.`); + } +} + export function assertProviderSelectedContext( context: Context, stepName: string, ): asserts context is ProviderSelectedOnboardFlowContext { - if (!context.model || !context.provider || !context.sandboxGpuConfig) { + assertProviderModelSelectedContext(context, stepName); + if (!context.sandboxGpuConfig) { throw new Error(`Onboarding state is incomplete before ${stepName}.`); } } diff --git a/src/lib/onboard/machine/flow-phases/provider-sandbox.test.ts b/src/lib/onboard/machine/flow-phases/provider-sandbox.test.ts index 1431a4bebf8..1589a2e694f 100644 --- a/src/lib/onboard/machine/flow-phases/provider-sandbox.test.ts +++ b/src/lib/onboard/machine/flow-phases/provider-sandbox.test.ts @@ -1,14 +1,16 @@ // SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -import { describe, expect, it } from "vitest"; +import { describe, expect, it, vi } from "vitest"; import { createSession } from "../../../state/onboard-session"; -import { advanceTo, branchTo } from "../result"; import type { OnboardFlowContext } from "../flow-context"; +import { advanceTo, branchTo } from "../result"; import { createProviderInferencePhase, createSandboxPhase } from "./provider-sandbox"; -function context(): OnboardFlowContext { +function context( + patch: Partial> = {}, +): OnboardFlowContext { return { resume: false, fresh: false, @@ -32,6 +34,7 @@ function context(): OnboardFlowContext { gpu: null, sandboxGpuConfig: { mode: "0" }, gpuPassthrough: false, + ...patch, }; } @@ -39,12 +42,17 @@ describe("provider/sandbox flow phases", () => { it("maps provider inference context updates and ordered FSM results", async () => { const phase = createProviderInferencePhase(async () => ({ context: { + session: createSession(), sandboxName: "my-assistant", provider: "nvidia-prod", model: "model", endpointUrl: "https://example.com/v1", credentialEnv: "NVIDIA_INFERENCE_API_KEY", + hermesAuthMethod: null, + hermesToolGateways: [], preferredInferenceApi: "openai-responses", + nimContainer: null, + webSearchConfig: null, }, result: [advanceTo("inference"), advanceTo("sandbox")], })); @@ -67,14 +75,18 @@ describe("provider/sandbox flow phases", () => { }); const phase = createSandboxPhase(async () => ({ context: { + session: createSession(), sandboxName: "my-assistant", + webSearchConfig: null, selectedMessagingChannels: ["telegram"], webSearchSupported: true, }, result: branchResult, })); - const result = await phase.run(context()); + const result = await phase.run( + context({ model: "model", provider: "nvidia-prod", sandboxGpuConfig: { mode: "0" } }), + ); expect(phase.state).toBe("sandbox"); expect(result.context).toMatchObject({ @@ -84,4 +96,23 @@ describe("provider/sandbox flow phases", () => { }); expect(result.result).toEqual(branchResult); }); + + it("rejects sandbox phase execution before sandbox GPU config is selected", async () => { + const runSandbox = vi.fn(async () => ({ + context: { + session: createSession(), + sandboxName: "my-assistant", + webSearchConfig: null, + selectedMessagingChannels: [], + webSearchSupported: false, + }, + result: branchTo("openclaw"), + })); + const phase = createSandboxPhase(runSandbox); + + await expect( + phase.run(context({ model: "model", provider: "nvidia-prod", sandboxGpuConfig: null })), + ).rejects.toThrow(/Onboarding state is incomplete before sandbox setup\./); + expect(runSandbox).not.toHaveBeenCalled(); + }); }); diff --git a/src/lib/onboard/machine/flow-phases/provider-sandbox.ts b/src/lib/onboard/machine/flow-phases/provider-sandbox.ts index e57d3edba69..0e6a3df12da 100644 --- a/src/lib/onboard/machine/flow-phases/provider-sandbox.ts +++ b/src/lib/onboard/machine/flow-phases/provider-sandbox.ts @@ -1,19 +1,32 @@ // SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -import type { OnboardFlowContext, OnboardFlowPhaseResult } from "../flow-context"; -import { mergeOnboardFlowContext, onboardFlowPhaseResult } from "../flow-context"; +import type { + OnboardFlowContext, + OnboardFlowPhaseResult, + ProviderModelSelectedContextUpdate, + ProviderModelSelectedOnboardFlowContext, + SandboxCreatedContextUpdate, +} from "../flow-context"; +import { + assertProviderSelectedContext, + mergeProviderModelSelectedContext, + mergeSandboxCreatedContext, + onboardFlowPhaseResult, +} from "../flow-context"; import type { OnboardSequencePhase } from "../sequence-runner"; type ProviderInferencePhaseHandler = ( context: Context, ) => Promise<{ - context: Partial; + context: ProviderModelSelectedContextUpdate; result: OnboardFlowPhaseResult["result"]; }>; -type SandboxPhaseHandler = (context: Context) => Promise<{ - context: Partial; +type SandboxPhaseHandler = ( + context: ProviderModelSelectedOnboardFlowContext, +) => Promise<{ + context: SandboxCreatedContextUpdate; result: OnboardFlowPhaseResult["result"]; }>; @@ -25,7 +38,7 @@ export function createProviderInferencePhase async run(context) { const result = await runProviderInference(context); return onboardFlowPhaseResult( - mergeOnboardFlowContext(context, result.context), + mergeProviderModelSelectedContext(context, result.context), result.result, ); }, @@ -38,9 +51,10 @@ export function createSandboxPhase( return { state: "sandbox", async run(context) { + assertProviderSelectedContext(context, "sandbox setup"); const result = await runSandbox(context); return onboardFlowPhaseResult( - mergeOnboardFlowContext(context, result.context), + mergeSandboxCreatedContext(context, result.context), result.result, ); }, diff --git a/src/lib/onboard/machine/flow-sequence.test.ts b/src/lib/onboard/machine/flow-sequence.test.ts index 9d0c30c88a1..51e850cd479 100644 --- a/src/lib/onboard/machine/flow-sequence.test.ts +++ b/src/lib/onboard/machine/flow-sequence.test.ts @@ -8,9 +8,9 @@ import { filterSafeUpdates, MACHINE_SNAPSHOT_VERSION, normalizeSession, - sanitizeFailure, type Session, type SessionUpdates, + sanitizeFailure, } from "../../state/onboard-session"; import type { OnboardFlowContext, OnboardFlowPhaseResult } from "./flow-context"; import { onboardFlowPhaseResult } from "./flow-context"; @@ -21,7 +21,7 @@ import { runOnboardSequenceWithRunner } from "./sequence-runner"; type Context = OnboardFlowContext; -function context(): Context { +function context(patch: Partial = {}): Context { return { resume: false, fresh: false, @@ -45,6 +45,7 @@ function context(): Context { gpu: null, sandboxGpuConfig: { mode: "0" }, gpuPassthrough: false, + ...patch, }; } @@ -141,8 +142,10 @@ describe("onboard flow phase sequence", () => { preflight: async (ctx) => result({ ...ctx, gpu: { type: "nvidia" }, gpuPassthrough: true }, "gateway"), gateway: async (ctx) => result(ctx, "provider_selection"), - providerInference: async (ctx) => result(ctx, "sandbox"), - sandbox: async (ctx) => onboardFlowPhaseResult(ctx, branchTo("openclaw")), + providerInference: async (ctx) => + result({ ...ctx, provider: "nvidia", model: "model" }, "sandbox"), + sandbox: async (ctx) => + onboardFlowPhaseResult({ ...ctx, sandboxName: "my-assistant" }, branchTo("openclaw")), openclaw: async (ctx) => result(ctx, "policies"), agentSetup: async (ctx) => result(ctx, "policies"), policies: async (ctx) => result(ctx, "finalizing"), @@ -156,6 +159,47 @@ describe("onboard flow phase sequence", () => { expect(preflight.result).toMatchObject({ next: "gateway" }); }); + it("rejects provider inference results that omit provider or model", async () => { + const phases = buildOnboardFlowPhaseSequence({ + preflight: async (ctx) => result(ctx, "gateway"), + gateway: async (ctx) => result(ctx, "provider_selection"), + providerInference: async (ctx) => + result({ ...ctx, model: "model", provider: null }, "sandbox"), + sandbox: async (ctx) => onboardFlowPhaseResult(ctx, branchTo("openclaw")), + openclaw: async (ctx) => result(ctx, "policies"), + agentSetup: async (ctx) => result(ctx, "policies"), + policies: async (ctx) => result(ctx, "finalizing"), + finalization: async (ctx) => result(ctx, "post_verify"), + postVerify: async (ctx) => onboardFlowPhaseResult(ctx, completeOnboardMachine()), + }); + + await expect(phases[2].run(context())).rejects.toThrow( + /Onboarding state is incomplete before provider inference result\./, + ); + }); + + it("rejects sandbox results that omit sandbox name", async () => { + const phases = buildOnboardFlowPhaseSequence({ + preflight: async (ctx) => result(ctx, "gateway"), + gateway: async (ctx) => result(ctx, "provider_selection"), + providerInference: async (ctx) => + result({ ...ctx, provider: "nvidia", model: "model" }, "sandbox"), + sandbox: async (ctx) => + onboardFlowPhaseResult({ ...ctx, sandboxName: null }, branchTo("openclaw")), + openclaw: async (ctx) => result(ctx, "policies"), + agentSetup: async (ctx) => result(ctx, "policies"), + policies: async (ctx) => result(ctx, "finalizing"), + finalization: async (ctx) => result(ctx, "post_verify"), + postVerify: async (ctx) => onboardFlowPhaseResult(ctx, completeOnboardMachine()), + }); + + await expect( + phases[3].run( + context({ provider: "nvidia", model: "model", sandboxGpuConfig: { mode: "0" } }), + ), + ).rejects.toThrow(/Onboarding state is incomplete before sandbox result\./); + }); + it("runs ordered provider results through runtime transition validation", async () => { const initialSession = createSession({ machine: { diff --git a/src/lib/onboard/machine/flow-sequence.ts b/src/lib/onboard/machine/flow-sequence.ts index d3cf2d1b560..0946af47509 100644 --- a/src/lib/onboard/machine/flow-sequence.ts +++ b/src/lib/onboard/machine/flow-sequence.ts @@ -2,6 +2,7 @@ // SPDX-License-Identifier: Apache-2.0 import type { OnboardFlowContext, OnboardFlowPhaseResult } from "./flow-context"; +import { assertProviderModelSelectedContext, assertSandboxCreatedContext } from "./flow-context"; import { createAgentSetupPhase, createFinalizationPhase, @@ -45,8 +46,40 @@ export function buildOnboardFlowPhaseSequence handlers.providerInference(context)), - createSandboxPhase((context) => handlers.sandbox(context)), + createProviderInferencePhase(async (context) => { + const result = await handlers.providerInference(context); + assertProviderModelSelectedContext(result.context, "provider inference result"); + return { + context: { + session: result.context.session, + sandboxName: result.context.sandboxName, + model: result.context.model, + provider: result.context.provider, + endpointUrl: result.context.endpointUrl, + credentialEnv: result.context.credentialEnv, + hermesAuthMethod: result.context.hermesAuthMethod, + hermesToolGateways: result.context.hermesToolGateways, + preferredInferenceApi: result.context.preferredInferenceApi, + nimContainer: result.context.nimContainer, + webSearchConfig: result.context.webSearchConfig, + }, + result: result.result, + }; + }), + createSandboxPhase(async (context) => { + const result = await handlers.sandbox(context); + assertSandboxCreatedContext(result.context, "sandbox result"); + return { + context: { + session: result.context.session, + sandboxName: result.context.sandboxName, + webSearchConfig: result.context.webSearchConfig, + selectedMessagingChannels: result.context.selectedMessagingChannels, + webSearchSupported: result.context.webSearchSupported, + }, + result: result.result, + }; + }), createOpenclawSetupPhase((context) => handlers.openclaw(context)), createAgentSetupPhase((context) => handlers.agentSetup(context)), createPoliciesPhase((context) => handlers.policies(context)),