diff --git a/packages/opencode/src/kilocode/plan-followup.ts b/packages/opencode/src/kilocode/plan-followup.ts index ab11e366f0c..a95f9a69c85 100644 --- a/packages/opencode/src/kilocode/plan-followup.ts +++ b/packages/opencode/src/kilocode/plan-followup.ts @@ -2,6 +2,8 @@ import { Telemetry } from "@kilocode/kilo-telemetry" import { Agent } from "@/agent/agent" import { Bus } from "@/bus" import { TuiEvent } from "@/cli/cmd/tui/event" +import { Flag } from "@/flag/flag" +import { Global } from "@/global" import { Identifier } from "@/id/id" import { Provider } from "@/provider/provider" import { Question } from "@/question" @@ -10,6 +12,8 @@ import { LLM } from "@/session/llm" import { MessageV2 } from "@/session/message-v2" import { Todo } from "@/session/todo" import { Log } from "@/util/log" +import fs from "fs/promises" +import path from "path" function toText(item: MessageV2.WithParts): string { return item.parts @@ -111,6 +115,41 @@ export namespace PlanFollowup { export const ANSWER_NEW_SESSION = "Start new session" export const ANSWER_CONTINUE = "Continue here" + async function resolveCodeModel(model: MessageV2.User["model"]) { + const saved = + Flag.KILO_CLIENT === "cli" + ? await fs + .readFile(path.join(Global.Path.state, "model.json"), "utf-8") + .then( + (item) => + JSON.parse(item) as { + model?: Record< + string, + { + providerID: string + modelID: string + } + > + }, + ) + .then((item) => item.model?.code) + .catch(() => undefined) + : undefined + if (saved) { + const match = await Provider.getModel(saved.providerID, saved.modelID).catch(() => undefined) + if (match) { + return { + providerID: saved.providerID, + modelID: saved.modelID, + } + } + } + + const agent = await Agent.get("code") + if (agent?.model) return agent.model + return model + } + async function resolvePlan(input: { assistant?: MessageV2.WithParts messages: MessageV2.WithParts[] @@ -211,6 +250,7 @@ export namespace PlanFollowup { model: MessageV2.User["model"] abort?: AbortSignal }) { + const model = await resolveCodeModel(input.model) const [handover, todos] = await Promise.all([ generateHandover({ messages: input.messages, model: input.model, abort: input.abort }), Todo.get(input.sessionID), @@ -231,7 +271,7 @@ export namespace PlanFollowup { await inject({ sessionID: next.id, agent: "code", - model: input.model, + model, text: sections.join("\n\n"), synthetic: false, }) @@ -289,10 +329,11 @@ export namespace PlanFollowup { if (answer === ANSWER_CONTINUE) { Telemetry.trackPlanFollowup(input.sessionID, "continue") + const model = await resolveCodeModel(user.model) await inject({ sessionID: input.sessionID, agent: "code", - model: user.model, + model, text: "Implement the plan above.", }) return "continue" diff --git a/packages/opencode/test/kilocode/plan-followup.test.ts b/packages/opencode/test/kilocode/plan-followup.test.ts index 593b258906b..badc8c42ced 100644 --- a/packages/opencode/test/kilocode/plan-followup.test.ts +++ b/packages/opencode/test/kilocode/plan-followup.test.ts @@ -12,19 +12,42 @@ import { LLM } from "../../src/session/llm" import { MessageV2 } from "../../src/session/message-v2" import { SessionPrompt } from "../../src/session/prompt" import { Todo } from "../../src/session/todo" +import { Global } from "../../src/global" import { Log } from "../../src/util/log" +import path from "path" +import fs from "fs/promises" import { tmpdir } from "../fixture/fixture" Log.init({ print: false }) +process.env.KILO_CLIENT = "cli" const model = { providerID: "openai", modelID: "gpt-4", } +const saved = { + providerID: "openai", + modelID: "gpt-5", +} + +const config = { + providerID: "openai", + modelID: "gpt-4.1", +} + +const statePath = path.join(Global.Path.state, "model.json") + async function withInstance(fn: () => Promise) { await using tmp = await tmpdir({ git: true }) - await Instance.provide({ directory: tmp.path, fn }) + await fs.rm(statePath, { force: true }).catch(() => {}) + await Instance.provide({ + directory: tmp.path, + fn: async () => { + await fs.rm(statePath, { force: true }).catch(() => {}) + await fn() + }, + }) } async function seed(input: { @@ -126,6 +149,20 @@ async function sessions() { return Array.fromAsync(Session.list()) } +async function waitQuestion(sessionID: string) { + for (let i = 0; i < 50; i++) { + const list = await Question.list() + const item = list.find((q) => q.sessionID === sessionID) + if (item) return item + await Bun.sleep(10) + } +} + +async function writeState(input: { model?: Record }) { + await fs.mkdir(Global.Path.state, { recursive: true }) + await fs.writeFile(statePath, JSON.stringify(input)) +} + const fakeAgent: Agent.Info = { name: "compaction", mode: "subagent", @@ -171,15 +208,33 @@ describe("plan follow-up", () => { abort: AbortSignal.any([]), }) - const list = await Question.list() - expect(list).toHaveLength(1) - await Question.reject(list[0].id) + const item = await waitQuestion(seeded.sessionID) + expect(item).toBeDefined() + if (!item) return + await Question.reject(item.id) await expect(pending).resolves.toBe("break") })) test("ask - returns continue and creates code message on Continue here", () => withInstance(async () => { + const get = spyOn(Agent, "get").mockImplementation(async (name: string) => { + if (name === "code") { + return { + name: "code", + mode: "primary", + permission: [], + options: {}, + model: saved, + } as any + } + return undefined as any + }) + using _ = { + [Symbol.dispose]() { + get.mockRestore() + }, + } const seeded = await seed({ text: "1. Build\n2. Test" }) const pending = PlanFollowup.ask({ sessionID: seeded.sessionID, @@ -187,9 +242,11 @@ describe("plan follow-up", () => { abort: AbortSignal.any([]), }) - const list = await Question.list() + const item = await waitQuestion(seeded.sessionID) + expect(item).toBeDefined() + if (!item) return await Question.reply({ - requestID: list[0].id, + requestID: item.id, answers: [[PlanFollowup.ANSWER_CONTINUE]], }) @@ -199,6 +256,7 @@ describe("plan follow-up", () => { expect(user?.info.role).toBe("user") if (!user || user.info.role !== "user") return expect(user.info.agent).toBe("code") + expect(user.info.model).toEqual(saved) const part = user.parts.find((item) => item.type === "text") expect(part?.type).toBe("text") @@ -216,8 +274,11 @@ describe("plan follow-up", () => { abort: AbortSignal.any([]), }) + const item = await waitQuestion(seeded.sessionID) + expect(item).toBeDefined() + if (!item) return await Question.reply({ - requestID: (await Question.list())[0].id, + requestID: item.id, answers: [["Add rollback support too"]], }) @@ -237,6 +298,24 @@ describe("plan follow-up", () => { test("ask - creates a new session on Start new session with handover and todos", () => withInstance(async () => { + const get = spyOn(Agent, "get").mockImplementation(async (name: string) => { + if (name === "code") { + return { + name: "code", + mode: "primary", + permission: [], + options: {}, + model: saved, + } as any + } + if (name === "compaction") return fakeAgent as any + return undefined as any + }) + using _file = { + [Symbol.dispose]() { + get.mockRestore() + }, + } const loop = spyOn(SessionPrompt, "loop").mockResolvedValue({ info: { id: "msg_test", @@ -268,9 +347,19 @@ describe("plan follow-up", () => { }, parts: [], }) - using _mocks = mockHandoverDeps( - "## Discoveries\n\nFound REST endpoints in src/api.ts\n\n## Relevant Files\n\n- src/api.ts: REST endpoints\n- src/db.ts: Database layer", - ) + const modelSpy = spyOn(Provider, "getModel").mockResolvedValue(fakeModel) + const llmSpy = spyOn(LLM, "stream").mockResolvedValue({ + text: Promise.resolve( + "## Discoveries\n\nFound REST endpoints in src/api.ts\n\n## Relevant Files\n\n- src/api.ts: REST endpoints\n- src/db.ts: Database layer", + ), + } as any) + using _mocks = { + llmSpy, + [Symbol.dispose]() { + modelSpy.mockRestore() + llmSpy.mockRestore() + }, + } using _loop = { [Symbol.dispose]() { loop.mockRestore() @@ -300,8 +389,11 @@ describe("plan follow-up", () => { abort: AbortSignal.any([]), }) + const item = await waitQuestion(seeded.sessionID) + expect(item).toBeDefined() + if (!item) return await Question.reply({ - requestID: (await Question.list())[0].id, + requestID: item.id, answers: [[PlanFollowup.ANSWER_NEW_SESSION]], }) @@ -323,6 +415,7 @@ describe("plan follow-up", () => { expect(user?.info.role).toBe("user") if (!user || user.info.role !== "user") throw new Error("expected seeded user message") expect(user.info.agent).toBe("code") + expect(user.info.model).toEqual(saved) const part = user.parts.find((item) => item.type === "text") expect(part?.type).toBe("text") @@ -344,6 +437,85 @@ describe("plan follow-up", () => { SessionPrompt.cancel(newSessionID) })) + test("ask - falls back to configured code model when saved CLI code model is unavailable", () => + withInstance(async () => { + await writeState({ model: { code: { providerID: "missing", modelID: "ghost" } } }) + const get = spyOn(Agent, "get").mockImplementation(async (name: string) => { + if (name === "code") { + return { + name: "code", + mode: "primary", + permission: [], + options: {}, + model: config, + } as any + } + return undefined as any + }) + using _ = { + [Symbol.dispose]() { + get.mockRestore() + }, + } + const seeded = await seed({ text: "1. Build\n2. Test" }) + const pending = PlanFollowup.ask({ + sessionID: seeded.sessionID, + messages: seeded.messages, + abort: AbortSignal.any([]), + }) + + const item = await waitQuestion(seeded.sessionID) + expect(item).toBeDefined() + if (!item) return + await Question.reply({ + requestID: item.id, + answers: [[PlanFollowup.ANSWER_CONTINUE]], + }) + + await expect(pending).resolves.toBe("continue") + + const user = await latestUser(seeded.sessionID) + expect(user?.info.role).toBe("user") + if (!user || user.info.role !== "user") return + expect(user.info.agent).toBe("code") + expect(user.info.model).toEqual(config) + })) + + test("ask - falls back to planning model when no saved or configured code model exists", () => + withInstance(async () => { + const get = spyOn(Agent, "get").mockImplementation(async (name: string) => { + if (name === "code") return undefined as any + return undefined as any + }) + using _ = { + [Symbol.dispose]() { + get.mockRestore() + }, + } + const seeded = await seed({ text: "1. Build\n2. Test" }) + const pending = PlanFollowup.ask({ + sessionID: seeded.sessionID, + messages: seeded.messages, + abort: AbortSignal.any([]), + }) + + const item = await waitQuestion(seeded.sessionID) + expect(item).toBeDefined() + if (!item) return + await Question.reply({ + requestID: item.id, + answers: [[PlanFollowup.ANSWER_CONTINUE]], + }) + + await expect(pending).resolves.toBe("continue") + + const user = await latestUser(seeded.sessionID) + expect(user?.info.role).toBe("user") + if (!user || user.info.role !== "user") return + expect(user.info.agent).toBe("code") + expect(user.info.model).toEqual(model) + })) + test("ask - new session omits handover section when LLM returns empty", () => withInstance(async () => { const loop = spyOn(SessionPrompt, "loop").mockResolvedValue({ @@ -387,8 +559,11 @@ describe("plan follow-up", () => { abort: AbortSignal.any([]), }) + const item = await waitQuestion(seeded.sessionID) + expect(item).toBeDefined() + if (!item) return await Question.reply({ - requestID: (await Question.list())[0].id, + requestID: item.id, answers: [[PlanFollowup.ANSWER_NEW_SESSION]], }) @@ -444,8 +619,9 @@ describe("plan follow-up", () => { abort: abort.signal, }) - const list = await Question.list() - expect(list).toHaveLength(1) + const item = await waitQuestion(seeded.sessionID) + expect(item).toBeDefined() + if (!item) return abort.abort() @@ -462,8 +638,11 @@ describe("plan follow-up", () => { abort: AbortSignal.any([]), }) + const item = await waitQuestion(seeded.sessionID) + expect(item).toBeDefined() + if (!item) return await Question.reply({ - requestID: (await Question.list())[0].id, + requestID: item.id, answers: [[" "]], })