diff --git a/src/lib/onboard/machine/core-flow-phases.ts b/src/lib/onboard/machine/core-flow-phases.ts index 91011a26a8d..af4a72b7734 100644 --- a/src/lib/onboard/machine/core-flow-phases.ts +++ b/src/lib/onboard/machine/core-flow-phases.ts @@ -2,7 +2,12 @@ // SPDX-License-Identifier: Apache-2.0 import type { WebSearchConfig } from "../../inference/web-search"; -import { assertProviderSelectedContext, type OnboardFlowContext } from "./flow-context"; +import { + assertProviderSelectedContext, + mergeProviderModelSelectedContext, + mergeSandboxCreatedContext, + type OnboardFlowContext, +} from "./flow-context"; import { runCoreOnboardFlowSequence } from "./flow-slices"; import { handleProviderInferenceState, @@ -75,8 +80,7 @@ export function createCoreOnboardFlowPhases< }); return { - context: { - ...context, + context: mergeProviderModelSelectedContext(context, { session: providerInferenceResult.session, sandboxName: providerInferenceResult.sandboxName, model: providerInferenceResult.model, @@ -88,7 +92,7 @@ export function createCoreOnboardFlowPhases< preferredInferenceApi: providerInferenceResult.preferredInferenceApi, nimContainer: providerInferenceResult.nimContainer, webSearchConfig: providerInferenceResult.webSearchConfig, - }, + }), result: providerInferenceResult.stateResults, }; }, @@ -121,14 +125,13 @@ export function createCoreOnboardFlowPhases< }); return { - context: { - ...context, + context: mergeSandboxCreatedContext(context, { session: sandboxStateResult.session, sandboxName: sandboxStateResult.sandboxName, - webSearchConfig: sandboxStateResult.webSearchConfig ?? null, + webSearchConfig: sandboxStateResult.webSearchConfig, selectedMessagingChannels: sandboxStateResult.selectedMessagingChannels, webSearchSupported: sandboxStateResult.webSearchSupported, - }, + }), result: sandboxStateResult.stateResult, }; }, diff --git a/src/lib/onboard/machine/flow-context.test.ts b/src/lib/onboard/machine/flow-context.test.ts index dbd7eafcdd0..8ddd580d298 100644 --- a/src/lib/onboard/machine/flow-context.test.ts +++ b/src/lib/onboard/machine/flow-context.test.ts @@ -8,6 +8,8 @@ import { assertProviderSelectedContext, assertSandboxCreatedContext, mergeOnboardFlowContext, + mergeProviderModelSelectedContext, + mergeSandboxCreatedContext, type OnboardFlowContext, onboardFlowPhaseResult, } from "./flow-context"; @@ -64,6 +66,29 @@ describe("onboard flow context helpers", () => { expect(result.result).toMatchObject({ next: "gateway", transitionKind: "advance" }); }); + it("merges provider/model-selected context updates", () => { + const context = mergeProviderModelSelectedContext(baseContext(), { + session: createSession(), + sandboxName: "my-assistant", + provider: "nvidia-prod", + model: "model", + endpointUrl: "https://example.test/v1", + credentialEnv: "NVIDIA_INFERENCE_API_KEY", + hermesAuthMethod: null, + hermesToolGateways: [], + preferredInferenceApi: "openai-responses", + nimContainer: null, + webSearchConfig: null, + }); + + expect(context).toMatchObject({ + sandboxName: "my-assistant", + provider: "nvidia-prod", + model: "model", + credentialEnv: "NVIDIA_INFERENCE_API_KEY", + }); + }); + it("asserts provider-selected context before sandbox setup", () => { const context = mergeOnboardFlowContext(baseContext(), { provider: "nvidia-prod", @@ -79,6 +104,35 @@ describe("onboard flow context helpers", () => { ); }); + it("merges sandbox-created context updates", () => { + const providerContext = mergeProviderModelSelectedContext(baseContext(), { + session: createSession(), + sandboxName: null, + provider: "nvidia-prod", + model: "model", + endpointUrl: null, + credentialEnv: null, + hermesAuthMethod: null, + hermesToolGateways: [], + preferredInferenceApi: null, + nimContainer: null, + webSearchConfig: null, + }); + const context = mergeSandboxCreatedContext(providerContext, { + session: createSession(), + sandboxName: "my-assistant", + webSearchConfig: null, + selectedMessagingChannels: ["telegram"], + webSearchSupported: true, + }); + + expect(context).toMatchObject({ + sandboxName: "my-assistant", + selectedMessagingChannels: ["telegram"], + webSearchSupported: true, + }); + }); + it("asserts sandbox-created context before final phases", () => { const context = mergeOnboardFlowContext(baseContext(), { sandboxName: "my-assistant", diff --git a/src/lib/onboard/machine/flow-context.ts b/src/lib/onboard/machine/flow-context.ts index 761e71681bb..bb7ace415ab 100644 --- a/src/lib/onboard/machine/flow-context.ts +++ b/src/lib/onboard/machine/flow-context.ts @@ -30,11 +30,16 @@ export interface OnboardFlowContext = Context & { - model: string; - provider: string; - sandboxGpuConfig: NonNullable; -}; +export type ProviderModelSelectedOnboardFlowContext = + Context & { + model: string; + provider: string; + }; + +export type ProviderSelectedOnboardFlowContext = + ProviderModelSelectedOnboardFlowContext & { + sandboxGpuConfig: NonNullable; + }; export type SandboxCreatedOnboardFlowContext = Context & { sandboxName: string; @@ -50,6 +55,28 @@ export interface OnboardFlowPhaseResult( context: Context, stepName: string, @@ -75,6 +102,20 @@ export function mergeOnboardFlowContext( return { ...context, ...patch }; } +export function mergeProviderModelSelectedContext( + context: Context, + patch: ProviderModelSelectedContextUpdate, +): ProviderModelSelectedOnboardFlowContext { + return { ...context, ...patch }; +} + +export function mergeSandboxCreatedContext( + context: ProviderModelSelectedOnboardFlowContext, + patch: SandboxCreatedContextUpdate, +): SandboxCreatedOnboardFlowContext { + return { ...context, ...patch }; +} + export function onboardFlowPhaseResult( context: Context, result: OnboardStateHandlerResult,