fix(cli): use configured Code model when implementing a plan

This commit is contained in:
Josh Holmer
2026-03-19 10:44:56 -04:00
parent 6ed7fdb40d
commit 1fc0e644dd
2 changed files with 237 additions and 17 deletions
@@ -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"
@@ -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<void>) {
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<string, { providerID: string; modelID: string }> }) {
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: [[" "]],
})