diff --git a/packages/tui/test/mini/fixture/footer-api.ts b/packages/tui/test/mini/fixture/footer-api.ts index bb1a05fb46d..83c06ca50fc 100644 --- a/packages/tui/test/mini/fixture/footer-api.ts +++ b/packages/tui/test/mini/fixture/footer-api.ts @@ -3,6 +3,10 @@ import type { FooterApi, FooterEvent, RunPrompt, StreamCommit } from "../../../s export function createFooterApiFixture(input: { events?: FooterEvent[]; commits?: StreamCommit[] } = {}) { const prompts = new Set<(input: RunPrompt) => void>() const closes = new Set<() => void>() + let ready!: () => void + const promptReady = new Promise((resolve) => { + ready = resolve + }) const events = input.events ?? [] const commits = input.commits ?? [] const calls: Array<{ type: "event"; value: FooterEvent } | { type: "commit"; value: StreamCommit }> = [] @@ -14,6 +18,7 @@ export function createFooterApiFixture(input: { events?: FooterEvent[]; commits? }, onPrompt(fn) { prompts.add(fn) + ready() return () => prompts.delete(fn) }, onClose(fn) { @@ -50,9 +55,12 @@ export function createFooterApiFixture(input: { events?: FooterEvent[]; commits? events, commits, calls, + promptReady, submit(text: string, mode?: RunPrompt["mode"]) { + if (prompts.size === 0) return false const prompt: RunPrompt = mode ? { text, parts: [], mode } : { text, parts: [] } for (const fn of [...prompts]) fn(prompt) + return true }, } } diff --git a/packages/tui/test/mini/runtime.test.ts b/packages/tui/test/mini/runtime.test.ts index 3b824cecd19..d912e600081 100644 --- a/packages/tui/test/mini/runtime.test.ts +++ b/packages/tui/test/mini/runtime.test.ts @@ -55,6 +55,9 @@ describe("run interactive runtime", () => { const api = ui.api const selected = defer>>() const catalogLoaded = defer() + const defaultModelReloaded = defer() + const modelShown = defer() + const turnStarted = defer() const model = catalogModel({ id: "resolved", providerID: "test", @@ -69,7 +72,17 @@ describe("run interactive runtime", () => { providers: [catalogProvider("test", "Test Provider")], models: [model], }) - const defaultModel = spyOn(sdk.model, "default").mockImplementation(() => selected.promise) + let defaultModelCalls = 0 + const defaultModel = spyOn(sdk.model, "default").mockImplementation(() => { + defaultModelCalls++ + if (defaultModelCalls === 2) defaultModelReloaded.resolve() + return selected.promise + }) + const emit = api.event.bind(api) + api.event = (event) => { + emit(event) + if (event.type === "model") modelShown.resolve() + } const task = runInteractiveDeferredMode( { @@ -110,6 +123,7 @@ describe("run interactive runtime", () => { runPromptTurn: async (input) => { turnAgent = input.agent turnModel = input.model + turnStarted.resolve() api.close() }, queuePromptTurn: async () => {}, @@ -133,8 +147,8 @@ describe("run interactive runtime", () => { location: { directory: "/tmp", project: { id: "pro-1", directory: "/tmp", canonical: "/tmp" } }, data: model, }) - while (defaultModel.mock.calls.length < 2) await Bun.sleep(0) - while (!events.some((event) => event.type === "model")) await Bun.sleep(0) + await defaultModelReloaded.promise + await modelShown.promise expect(events).toContainEqual({ type: "model", model: "Resolved Model ยท Test Provider", @@ -142,8 +156,9 @@ describe("run interactive runtime", () => { }) expect(lifecycle.onCycleVariant?.()).toMatchObject({ status: "variant low", variant: "low" }) lifecycle.onAgentSelect?.("review") - ui.submit("hello") - while (!turnModel) await Bun.sleep(0) + await ui.promptReady + expect(ui.submit("hello")).toBe(true) + await turnStarted.promise expect(turnAgent).toBe("review") expect(turnModel).toEqual({ providerID: "test", modelID: "resolved" }) await task