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
58 changes: 32 additions & 26 deletions src/lib/onboard/machine/core-flow-phases.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -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 () => {
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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" },
});
});

Expand Down
155 changes: 74 additions & 81 deletions src/lib/onboard/machine/core-flow-phases.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -52,91 +52,84 @@ export function createCoreOnboardFlowPhases<
>(
options: CoreOnboardFlowPhaseOptions<Context, Host, MessagingChannelConfig, ResourceProfile>,
): [OnboardSequencePhase<Context>, OnboardSequencePhase<Context>] {
const providerInferencePhase: OnboardSequencePhase<Context> = {
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<Context> = {
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<Context>(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<Context>(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];
}
Expand Down
71 changes: 44 additions & 27 deletions src/lib/onboard/machine/flow-phases/provider-sandbox.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -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,
Expand All @@ -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,
Expand All @@ -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();
});
});
Loading
Loading