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
19 changes: 11 additions & 8 deletions src/lib/onboard/machine/core-flow-phases.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -75,8 +80,7 @@ export function createCoreOnboardFlowPhases<
});

return {
context: {
...context,
context: mergeProviderModelSelectedContext(context, {
session: providerInferenceResult.session,
sandboxName: providerInferenceResult.sandboxName,
model: providerInferenceResult.model,
Expand All @@ -88,7 +92,7 @@ export function createCoreOnboardFlowPhases<
preferredInferenceApi: providerInferenceResult.preferredInferenceApi,
nimContainer: providerInferenceResult.nimContainer,
webSearchConfig: providerInferenceResult.webSearchConfig,
},
}),
result: providerInferenceResult.stateResults,
};
},
Expand Down Expand Up @@ -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,
};
},
Expand Down
54 changes: 54 additions & 0 deletions src/lib/onboard/machine/flow-context.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,8 @@ import {
assertProviderSelectedContext,
assertSandboxCreatedContext,
mergeOnboardFlowContext,
mergeProviderModelSelectedContext,
mergeSandboxCreatedContext,
type OnboardFlowContext,
onboardFlowPhaseResult,
} from "./flow-context";
Expand Down Expand Up @@ -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",
Expand All @@ -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",
Expand Down
51 changes: 46 additions & 5 deletions src/lib/onboard/machine/flow-context.ts
Original file line number Diff line number Diff line change
Expand Up @@ -30,11 +30,16 @@ export interface OnboardFlowContext<Agent = unknown, Gpu = unknown, SandboxGpuCo
gpuPassthrough: boolean;
}

export type ProviderSelectedOnboardFlowContext<Context extends OnboardFlowContext> = Context & {
model: string;
provider: string;
sandboxGpuConfig: NonNullable<Context["sandboxGpuConfig"]>;
};
export type ProviderModelSelectedOnboardFlowContext<Context extends OnboardFlowContext> =
Context & {
model: string;
provider: string;
};

export type ProviderSelectedOnboardFlowContext<Context extends OnboardFlowContext> =
ProviderModelSelectedOnboardFlowContext<Context> & {
sandboxGpuConfig: NonNullable<Context["sandboxGpuConfig"]>;
};

export type SandboxCreatedOnboardFlowContext<Context extends OnboardFlowContext> = Context & {
sandboxName: string;
Expand All @@ -50,6 +55,28 @@ export interface OnboardFlowPhaseResult<Context extends OnboardFlowContext = Onb
result: OnboardStateHandlerResult;
}

export interface ProviderModelSelectedContextUpdate {
session: Session | null;
sandboxName: string | null;
model: string;
provider: string;
endpointUrl: string | null;
credentialEnv: string | null;
hermesAuthMethod: string | null;
hermesToolGateways: string[];
preferredInferenceApi: string | null;
nimContainer: string | null;
webSearchConfig: WebSearchConfig | null;
}

export interface SandboxCreatedContextUpdate {
session: Session | null;
sandboxName: string;
webSearchConfig: WebSearchConfig | null;
selectedMessagingChannels: string[];
webSearchSupported: boolean;
}

export function assertProviderSelectedContext<Context extends OnboardFlowContext>(
context: Context,
stepName: string,
Expand All @@ -75,6 +102,20 @@ export function mergeOnboardFlowContext<Context extends OnboardFlowContext>(
return { ...context, ...patch };
}

export function mergeProviderModelSelectedContext<Context extends OnboardFlowContext>(
context: Context,
patch: ProviderModelSelectedContextUpdate,
): ProviderModelSelectedOnboardFlowContext<Context> {
return { ...context, ...patch };
}

export function mergeSandboxCreatedContext<Context extends OnboardFlowContext>(
context: ProviderModelSelectedOnboardFlowContext<Context>,
patch: SandboxCreatedContextUpdate,
): SandboxCreatedOnboardFlowContext<Context> {
return { ...context, ...patch };
}

export function onboardFlowPhaseResult<Context extends OnboardFlowContext>(
context: Context,
result: OnboardStateHandlerResult,
Expand Down