diff --git a/src/lib/onboard/machine/core-flow-phases.test.ts b/src/lib/onboard/machine/core-flow-phases.test.ts index 88ecd68056b..9bf53e9b162 100644 --- a/src/lib/onboard/machine/core-flow-phases.test.ts +++ b/src/lib/onboard/machine/core-flow-phases.test.ts @@ -185,12 +185,12 @@ function createPhases( } describe("core onboard flow phases", () => { - it("runs provider selection and carries inference output into the flow context", async () => { - const [providerPhase] = createPhases(); + it("carries provider selection output into sandbox setup", async () => { + const [providerPhase, sandboxPhase] = createPhases(); - const result = await providerPhase.run(context()); + const providerResult = await providerPhase.run(context()); - expect(result.context).toMatchObject({ + expect(providerResult.context).toMatchObject({ sandboxName: "my-sandbox", model: "nvidia/test", provider: "nim", @@ -200,7 +200,26 @@ describe("core onboard flow phases", () => { preferredInferenceApi: "chat", nimContainer: "nim-test", }); - expect(Array.isArray(result.result)).toBe(true); + expect(Array.isArray(providerResult.result)).toBe(true); + + const sandboxResult = await sandboxPhase.run(providerResult.context); + + expect(sandboxResult.context).toMatchObject({ + sandboxName: "created-sandbox", + model: "nvidia/test", + provider: "nim", + endpointUrl: "https://example.test/v1", + credentialEnv: "NVIDIA_INFERENCE_API_KEY", + fromDockerfile: null, + gpu: { platform: "linux" }, + sandboxGpuConfig: { mode: "cdi" }, + gpuPassthrough: true, + hermesToolGateways: ["local"], + preferredInferenceApi: "chat", + nimContainer: "nim-test", + selectedMessagingChannels: ["slack", "discord"], + webSearchSupported: true, + }); }); it("passes fresh context through to provider setup recovery policy", async () => { @@ -228,7 +247,7 @@ describe("core onboard flow phases", () => { it("uses normalized context Hermes tool gateways for provider inference resume", async () => { const setupInference = vi.fn(async () => ({ ok: true as const })); - const [providerPhase] = createPhases({ + const [providerPhase, sandboxPhase] = createPhases({ providerDeps: { ensureResumeProviderReady: vi.fn(async () => ({ forceInferenceSetup: false, @@ -272,29 +291,16 @@ describe("core onboard flow phases", () => { { allowToolsIncompatible: false }, ); expect(result.context.hermesToolGateways).toEqual(["nous-web"]); - }); - - it("runs sandbox setup only after provider state is complete", async () => { - const [, sandboxPhase] = createPhases(); - await expect(sandboxPhase.run(context())).rejects.toThrow( - "Onboarding state is incomplete before sandbox setup.", - ); + const sandboxResult = await sandboxPhase.run(result.context); - const result = await sandboxPhase.run( - context({ - model: "nvidia/test", - provider: "nim", - hermesToolGateways: ["local"], - preferredInferenceApi: "chat", - nimContainer: "nim-test", - }), - ); - - expect(result.context).toMatchObject({ + expect(sandboxResult.context).toMatchObject({ sandboxName: "created-sandbox", - selectedMessagingChannels: ["slack", "discord"], - webSearchSupported: true, + model: "nvidia/test", + provider: "hermes", + credentialEnv: "HERMES_API_KEY", + hermesToolGateways: ["nous-web"], + sandboxGpuConfig: { mode: "cdi" }, }); }); diff --git a/src/lib/onboard/machine/core-flow-phases.ts b/src/lib/onboard/machine/core-flow-phases.ts index 86029aceeae..de183ad4f7a 100644 --- a/src/lib/onboard/machine/core-flow-phases.ts +++ b/src/lib/onboard/machine/core-flow-phases.ts @@ -3,11 +3,11 @@ import type { WebSearchConfig } from "../../inference/web-search"; import { - assertProviderSelectedContext, mergeProviderModelSelectedContext, mergeSandboxCreatedContext, type OnboardFlowContext, } from "./flow-context"; +import { createProviderInferencePhase, createSandboxPhase } from "./flow-phases/provider-sandbox"; import { runCoreOnboardFlowSequence } from "./flow-slices"; import { handleProviderInferenceState, @@ -52,91 +52,84 @@ export function createCoreOnboardFlowPhases< >( options: CoreOnboardFlowPhaseOptions, ): [OnboardSequencePhase, OnboardSequencePhase] { - const providerInferencePhase: OnboardSequencePhase = { - state: "provider_selection", - async run(context) { - const providerInferenceResult = await handleProviderInferenceState({ - resume: context.resume, - fresh: context.fresh, - session: context.session, - gpu: context.gpu, - sandboxName: context.sandboxName, - agent: context.agent, - forceProviderSelection: options.forceProviderSelection, - initial: { - model: context.model, - provider: context.provider, - endpointUrl: context.endpointUrl, - credentialEnv: context.credentialEnv, - hermesAuthMethod: context.hermesAuthMethod, - hermesToolGateways: context.hermesToolGateways, - preferredInferenceApi: context.preferredInferenceApi, - nimContainer: context.nimContainer, - webSearchConfig: context.webSearchConfig, - }, - selectedMessagingChannels: context.selectedMessagingChannels, - env: options.env, - constants: options.constants, - deps: options.providerDeps, - }); - - return { - context: mergeProviderModelSelectedContext(context, { - session: providerInferenceResult.session, - sandboxName: providerInferenceResult.sandboxName, - model: providerInferenceResult.model, - provider: providerInferenceResult.provider, - endpointUrl: providerInferenceResult.endpointUrl, - credentialEnv: providerInferenceResult.credentialEnv, - hermesAuthMethod: providerInferenceResult.hermesAuthMethod, - hermesToolGateways: providerInferenceResult.hermesToolGateways, - preferredInferenceApi: providerInferenceResult.preferredInferenceApi, - nimContainer: providerInferenceResult.nimContainer, - webSearchConfig: providerInferenceResult.webSearchConfig, - }), - result: providerInferenceResult.stateResults, - }; - }, - }; - - const sandboxPhase: OnboardSequencePhase = { - state: "sandbox", - async run(context) { - assertProviderSelectedContext(context, "sandbox setup"); - const sandboxStateResult = await handleSandboxState({ - resume: context.resume, - fresh: context.fresh, - resumeAgentChanged: options.sandbox.resumeAgentChanged, - session: context.session, - sandboxName: context.sandboxName, + const providerInferencePhase = createProviderInferencePhase(async (context) => { + const providerInferenceResult = await handleProviderInferenceState({ + resume: context.resume, + fresh: context.fresh, + session: context.session, + gpu: context.gpu, + sandboxName: context.sandboxName, + agent: context.agent, + forceProviderSelection: options.forceProviderSelection, + initial: { model: context.model, provider: context.provider, + endpointUrl: context.endpointUrl, + credentialEnv: context.credentialEnv, + hermesAuthMethod: context.hermesAuthMethod, + hermesToolGateways: context.hermesToolGateways, + preferredInferenceApi: context.preferredInferenceApi, nimContainer: context.nimContainer, webSearchConfig: context.webSearchConfig, - selectedMessagingChannels: context.selectedMessagingChannels, - fromDockerfile: context.fromDockerfile, - agent: context.agent, - gpu: context.gpu, - preferredInferenceApi: context.preferredInferenceApi, - sandboxGpuConfig: context.sandboxGpuConfig, - hermesToolGateways: context.hermesToolGateways, - controlUiPort: options.sandbox.controlUiPort, - rootDir: options.sandbox.rootDir, - deps: options.sandboxDeps, - }); + }, + selectedMessagingChannels: context.selectedMessagingChannels, + env: options.env, + constants: options.constants, + deps: options.providerDeps, + }); - return { - context: mergeSandboxCreatedContext(context, { - session: sandboxStateResult.session, - sandboxName: sandboxStateResult.sandboxName, - webSearchConfig: sandboxStateResult.webSearchConfig, - selectedMessagingChannels: sandboxStateResult.selectedMessagingChannels, - webSearchSupported: sandboxStateResult.webSearchSupported, - }), - result: sandboxStateResult.stateResult, - }; - }, - }; + return { + context: mergeProviderModelSelectedContext(context, { + session: providerInferenceResult.session, + sandboxName: providerInferenceResult.sandboxName, + model: providerInferenceResult.model, + provider: providerInferenceResult.provider, + endpointUrl: providerInferenceResult.endpointUrl, + credentialEnv: providerInferenceResult.credentialEnv, + hermesAuthMethod: providerInferenceResult.hermesAuthMethod, + hermesToolGateways: providerInferenceResult.hermesToolGateways, + preferredInferenceApi: providerInferenceResult.preferredInferenceApi, + nimContainer: providerInferenceResult.nimContainer, + webSearchConfig: providerInferenceResult.webSearchConfig, + }), + result: providerInferenceResult.stateResults, + }; + }); + + const sandboxPhase = createSandboxPhase(async (context) => { + const sandboxStateResult = await handleSandboxState({ + resume: context.resume, + fresh: context.fresh, + resumeAgentChanged: options.sandbox.resumeAgentChanged, + session: context.session, + sandboxName: context.sandboxName, + model: context.model, + provider: context.provider, + nimContainer: context.nimContainer, + webSearchConfig: context.webSearchConfig, + selectedMessagingChannels: context.selectedMessagingChannels, + fromDockerfile: context.fromDockerfile, + agent: context.agent, + gpu: context.gpu, + preferredInferenceApi: context.preferredInferenceApi, + sandboxGpuConfig: context.sandboxGpuConfig, + hermesToolGateways: context.hermesToolGateways, + controlUiPort: options.sandbox.controlUiPort, + rootDir: options.sandbox.rootDir, + deps: options.sandboxDeps, + }); + + return { + context: mergeSandboxCreatedContext(context, { + session: sandboxStateResult.session, + sandboxName: sandboxStateResult.sandboxName, + webSearchConfig: sandboxStateResult.webSearchConfig, + selectedMessagingChannels: sandboxStateResult.selectedMessagingChannels, + webSearchSupported: sandboxStateResult.webSearchSupported, + }), + result: sandboxStateResult.stateResult, + }; + }); return [providerInferencePhase, sandboxPhase]; } 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 1589a2e694f..ede225a01a3 100644 --- a/src/lib/onboard/machine/flow-phases/provider-sandbox.test.ts +++ b/src/lib/onboard/machine/flow-phases/provider-sandbox.test.ts @@ -39,9 +39,10 @@ function context( } describe("provider/sandbox flow phases", () => { - it("maps provider inference context updates and ordered FSM results", async () => { - const phase = createProviderInferencePhase(async () => ({ + it("passes full context and results through the shared handoff", async () => { + const providerPhase = createProviderInferencePhase(async (current) => ({ context: { + ...current, session: createSession(), sandboxName: "my-assistant", provider: "nvidia-prod", @@ -56,25 +57,12 @@ describe("provider/sandbox flow phases", () => { }, result: [advanceTo("inference"), advanceTo("sandbox")], })); - - const result = await phase.run(context()); - - expect(phase.state).toBe("provider_selection"); - expect(result.context).toMatchObject({ - sandboxName: "my-assistant", - provider: "nvidia-prod", - model: "model", - preferredInferenceApi: "openai-responses", - }); - expect(result.result).toEqual([advanceTo("inference"), advanceTo("sandbox")]); - }); - - it("maps sandbox context updates and branch result", async () => { const branchResult = branchTo("openclaw", { metadata: { sandboxName: "my-assistant", state: "sandbox" }, }); - const phase = createSandboxPhase(async () => ({ + const runSandbox = vi.fn(async (current) => ({ context: { + ...current, session: createSession(), sandboxName: "my-assistant", webSearchConfig: null, @@ -83,23 +71,46 @@ describe("provider/sandbox flow phases", () => { }, result: branchResult, })); + const sandboxPhase = createSandboxPhase(runSandbox); - const result = await phase.run( - context({ model: "model", provider: "nvidia-prod", sandboxGpuConfig: { mode: "0" } }), + const providerResult = await providerPhase.run( + context({ fromDockerfile: "Dockerfile", selectedMessagingChannels: ["slack"] }), ); + const sandboxResult = await sandboxPhase.run(providerResult.context); - expect(phase.state).toBe("sandbox"); - expect(result.context).toMatchObject({ + expect(providerPhase.state).toBe("provider_selection"); + expect(sandboxPhase.state).toBe("sandbox"); + expect(runSandbox).toHaveBeenCalledWith( + expect.objectContaining({ + provider: "nvidia-prod", + model: "model", + credentialEnv: "NVIDIA_INFERENCE_API_KEY", + fromDockerfile: "Dockerfile", + sandboxGpuConfig: { mode: "0" }, + selectedMessagingChannels: ["slack"], + }), + ); + expect(sandboxResult.context).toMatchObject({ sandboxName: "my-assistant", + provider: "nvidia-prod", + model: "model", + credentialEnv: "NVIDIA_INFERENCE_API_KEY", + fromDockerfile: "Dockerfile", selectedMessagingChannels: ["telegram"], webSearchSupported: true, }); - expect(result.result).toEqual(branchResult); + expect(providerResult.result).toEqual([advanceTo("inference"), advanceTo("sandbox")]); + expect(sandboxResult.result).toEqual(branchResult); }); - it("rejects sandbox phase execution before sandbox GPU config is selected", async () => { - const runSandbox = vi.fn(async () => ({ + it.each([ + "model", + "provider", + "sandboxGpuConfig", + ] as const)("rejects sandbox phase execution before %s is selected (#5938)", async (missingField) => { + const runSandbox = vi.fn(async (current) => ({ context: { + ...current, session: createSession(), sandboxName: "my-assistant", webSearchConfig: null, @@ -109,10 +120,16 @@ describe("provider/sandbox flow phases", () => { result: branchTo("openclaw"), })); const phase = createSandboxPhase(runSandbox); + const incomplete = context({ + model: "model", + provider: "nvidia-prod", + sandboxGpuConfig: { mode: "0" }, + }); + incomplete[missingField] = null; - await expect( - phase.run(context({ model: "model", provider: "nvidia-prod", sandboxGpuConfig: null })), - ).rejects.toThrow(/Onboarding state is incomplete before sandbox setup\./); + await expect(phase.run(incomplete)).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 0e6a3df12da..739241adcbc 100644 --- a/src/lib/onboard/machine/flow-phases/provider-sandbox.ts +++ b/src/lib/onboard/machine/flow-phases/provider-sandbox.ts @@ -4,29 +4,24 @@ import type { OnboardFlowContext, OnboardFlowPhaseResult, - ProviderModelSelectedContextUpdate, ProviderModelSelectedOnboardFlowContext, - SandboxCreatedContextUpdate, -} from "../flow-context"; -import { - assertProviderSelectedContext, - mergeProviderModelSelectedContext, - mergeSandboxCreatedContext, - onboardFlowPhaseResult, + ProviderSelectedOnboardFlowContext, + SandboxCreatedOnboardFlowContext, } from "../flow-context"; +import { assertProviderSelectedContext, onboardFlowPhaseResult } from "../flow-context"; import type { OnboardSequencePhase } from "../sequence-runner"; type ProviderInferencePhaseHandler = ( context: Context, ) => Promise<{ - context: ProviderModelSelectedContextUpdate; + context: ProviderModelSelectedOnboardFlowContext; result: OnboardFlowPhaseResult["result"]; }>; type SandboxPhaseHandler = ( - context: ProviderModelSelectedOnboardFlowContext, + context: ProviderSelectedOnboardFlowContext, ) => Promise<{ - context: SandboxCreatedContextUpdate; + context: SandboxCreatedOnboardFlowContext; result: OnboardFlowPhaseResult["result"]; }>; @@ -37,10 +32,7 @@ export function createProviderInferencePhase state: "provider_selection", async run(context) { const result = await runProviderInference(context); - return onboardFlowPhaseResult( - mergeProviderModelSelectedContext(context, result.context), - result.result, - ); + return onboardFlowPhaseResult(result.context, result.result); }, }; } @@ -53,10 +45,7 @@ export function createSandboxPhase( async run(context) { assertProviderSelectedContext(context, "sandbox setup"); const result = await runSandbox(context); - return onboardFlowPhaseResult( - mergeSandboxCreatedContext(context, result.context), - result.result, - ); + return onboardFlowPhaseResult(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 51e850cd479..c4f45ba5059 100644 --- a/src/lib/onboard/machine/flow-sequence.test.ts +++ b/src/lib/onboard/machine/flow-sequence.test.ts @@ -229,7 +229,13 @@ describe("onboard flow phase sequence", () => { }); const run = await runOnboardSequenceWithRunner({ - context: context(), + context: context({ + credentialEnv: "NVIDIA_INFERENCE_API_KEY", + fromDockerfile: "Dockerfile", + hermesToolGateways: ["local"], + sandboxGpuConfig: { mode: "sentinel" }, + selectedMessagingChannels: ["slack"], + }), runtime: createRuntime(initialSession), phases, }); @@ -243,6 +249,13 @@ describe("onboard flow phase sequence", () => { provider: "nvidia", model: "model", sandboxName: "my-assistant", + credentialEnv: "NVIDIA_INFERENCE_API_KEY", + fromDockerfile: "Dockerfile", + gpu: { type: "nvidia" }, + sandboxGpuConfig: { mode: "sentinel" }, + gpuPassthrough: true, + hermesToolGateways: ["local"], + selectedMessagingChannels: ["slack"], }); }); }); diff --git a/src/lib/onboard/machine/flow-sequence.ts b/src/lib/onboard/machine/flow-sequence.ts index 0946af47509..5463d8cc827 100644 --- a/src/lib/onboard/machine/flow-sequence.ts +++ b/src/lib/onboard/machine/flow-sequence.ts @@ -47,38 +47,14 @@ export function buildOnboardFlowPhaseSequence { - 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, - }; + const { context: nextContext, result } = await handlers.providerInference(context); + assertProviderModelSelectedContext(nextContext, "provider inference result"); + return { context: nextContext, 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, - }; + const { context: nextContext, result } = await handlers.sandbox(context); + assertSandboxCreatedContext(nextContext, "sandbox result"); + return { context: nextContext, result }; }), createOpenclawSetupPhase((context) => handlers.openclaw(context)), createAgentSetupPhase((context) => handlers.agentSetup(context)),