mirror of
https://github.com/Kilo-Org/kilocode.git
synced 2026-09-24 16:02:55 +08:00
fix(cli): use configured Code model when implementing a plan
This commit is contained in:
@@ -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: [[" "]],
|
||||
})
|
||||
|
||||
|
||||
Reference in New Issue
Block a user