diff --git a/packages/opencode/src/kilocode/plan-followup.ts b/packages/opencode/src/kilocode/plan-followup.ts index a95f9a69c85..4dc71b89411 100644 --- a/packages/opencode/src/kilocode/plan-followup.ts +++ b/packages/opencode/src/kilocode/plan-followup.ts @@ -115,39 +115,55 @@ 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 = + function resolveVariant(input: { value: string | undefined; model: Provider.Model | undefined }) { + if (!input.value) return undefined + if (!input.model?.variants?.[input.value]) return undefined + return input.value + } + + async function resolveCodeModel(input: Pick) { + const state = 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 - } - > + model?: Record + variant?: Record }, ) - .then((item) => item.model?.code) .catch(() => undefined) : undefined + const saved = state?.model?.code if (saved) { - const match = await Provider.getModel(saved.providerID, saved.modelID).catch(() => undefined) - if (match) { + const full = await Provider.getModel(saved.providerID, saved.modelID).catch(() => undefined) + if (full) { + const key = `${saved.providerID}/${saved.modelID}` return { - providerID: saved.providerID, - modelID: saved.modelID, + model: saved, + variant: resolveVariant({ + value: state?.variant?.[key], + model: full, + }), } } } const agent = await Agent.get("code") - if (agent?.model) return agent.model - return model + if (agent?.model) { + const full = agent.variant + ? await Provider.getModel(agent.model.providerID, agent.model.modelID).catch(() => undefined) + : undefined + return { + model: agent.model, + variant: resolveVariant({ + value: agent.variant, + model: full, + }), + } + } + return input } async function resolvePlan(input: { @@ -180,6 +196,7 @@ export namespace PlanFollowup { sessionID: string agent: string model: MessageV2.User["model"] + variant?: MessageV2.User["variant"] text: string synthetic?: boolean }) { @@ -192,6 +209,7 @@ export namespace PlanFollowup { }, agent: input.agent, model: input.model, + variant: input.variant, } await Session.updateMessage(msg) await Session.updatePart({ @@ -248,9 +266,13 @@ export namespace PlanFollowup { plan: string messages: MessageV2.WithParts[] model: MessageV2.User["model"] + variant?: MessageV2.User["variant"] abort?: AbortSignal }) { - const model = await resolveCodeModel(input.model) + const code = await resolveCodeModel({ + model: input.model, + variant: input.variant, + }) const [handover, todos] = await Promise.all([ generateHandover({ messages: input.messages, model: input.model, abort: input.abort }), Todo.get(input.sessionID), @@ -271,7 +293,8 @@ export namespace PlanFollowup { await inject({ sessionID: next.id, agent: "code", - model, + model: code.model, + variant: code.variant, text: sections.join("\n\n"), synthetic: false, }) @@ -322,6 +345,7 @@ export namespace PlanFollowup { plan, messages: input.messages, model: user.model, + variant: user.variant, abort: input.abort, }) return "break" @@ -329,11 +353,15 @@ export namespace PlanFollowup { if (answer === ANSWER_CONTINUE) { Telemetry.trackPlanFollowup(input.sessionID, "continue") - const model = await resolveCodeModel(user.model) + const code = await resolveCodeModel({ + model: user.model, + variant: user.variant, + }) await inject({ sessionID: input.sessionID, agent: "code", - model, + model: code.model, + variant: code.variant, text: "Implement the plan above.", }) return "continue" @@ -344,6 +372,7 @@ export namespace PlanFollowup { sessionID: input.sessionID, agent: "plan", model: user.model, + variant: user.variant, text: answer, }) return "continue" diff --git a/packages/opencode/test/kilocode/plan-followup.test.ts b/packages/opencode/test/kilocode/plan-followup.test.ts index badc8c42ced..bc19cdbc5ce 100644 --- a/packages/opencode/test/kilocode/plan-followup.test.ts +++ b/packages/opencode/test/kilocode/plan-followup.test.ts @@ -31,12 +31,18 @@ const saved = { modelID: "gpt-5", } +const savedVar = "high" + const config = { providerID: "openai", modelID: "gpt-4.1", } +const configVar = "max" +const planVar = "medium" + const statePath = path.join(Global.Path.state, "model.json") +const savedKey = `${saved.providerID}/${saved.modelID}` async function withInstance(fn: () => Promise) { await using tmp = await tmpdir({ git: true }) @@ -52,6 +58,7 @@ async function withInstance(fn: () => Promise) { async function seed(input: { text: string + variant?: string tools?: Array<{ tool: string; input: Record; output: string }> }) { const session = await Session.create({}) @@ -64,6 +71,7 @@ async function seed(input: { }, agent: "plan", model, + variant: input.variant, }) await Session.updatePart({ id: Identifier.ascending("part"), @@ -158,7 +166,10 @@ async function waitQuestion(sessionID: string) { } } -async function writeState(input: { model?: Record }) { +async function writeState(input: { + model?: Record + variant?: Record +}) { await fs.mkdir(Global.Path.state, { recursive: true }) await fs.writeFile(statePath, JSON.stringify(input)) } @@ -178,6 +189,19 @@ const fakeModel = { capabilities: {}, } as Provider.Model +function full(input: { providerID: string; modelID: string }, vars: string[]) { + return { + ...fakeModel, + id: input.modelID, + providerID: input.providerID, + variants: Object.fromEntries(vars.map((item) => [item, {}])), + } as Provider.Model +} + +const savedFull = full(saved, [savedVar, "low"]) +const savedConfigFull = full(saved, [configVar, "low"]) +const configFull = full(config, [configVar, "low"]) + function mockHandoverDeps(text: string, opts?: { agent?: Agent.Info | null }) { const agentSpy = spyOn(Agent, "get").mockResolvedValue( (opts?.agent === null ? undefined : (opts?.agent ?? fakeAgent)) as any, @@ -226,13 +250,16 @@ describe("plan follow-up", () => { permission: [], options: {}, model: saved, + variant: configVar, } as any } return undefined as any }) + const modelSpy = spyOn(Provider, "getModel").mockResolvedValue(savedConfigFull) using _ = { [Symbol.dispose]() { get.mockRestore() + modelSpy.mockRestore() }, } const seeded = await seed({ text: "1. Build\n2. Test" }) @@ -257,6 +284,7 @@ describe("plan follow-up", () => { if (!user || user.info.role !== "user") return expect(user.info.agent).toBe("code") expect(user.info.model).toEqual(saved) + expect(user.info.variant).toBe(configVar) const part = user.parts.find((item) => item.type === "text") expect(part?.type).toBe("text") @@ -306,6 +334,7 @@ describe("plan follow-up", () => { permission: [], options: {}, model: saved, + variant: configVar, } as any } if (name === "compaction") return fakeAgent as any @@ -347,7 +376,10 @@ describe("plan follow-up", () => { }, parts: [], }) - const modelSpy = spyOn(Provider, "getModel").mockResolvedValue(fakeModel) + const modelSpy = spyOn(Provider, "getModel").mockImplementation(async (providerID: string, modelID: string) => { + if (providerID === saved.providerID && modelID === saved.modelID) return savedConfigFull + return 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", @@ -416,6 +448,7 @@ describe("plan follow-up", () => { 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) + expect(user.info.variant).toBe(configVar) const part = user.parts.find((item) => item.type === "text") expect(part?.type).toBe("text") @@ -437,6 +470,61 @@ describe("plan follow-up", () => { SessionPrompt.cancel(newSessionID) })) + test("ask - prefers saved code variant over configured code variant", () => + withInstance(async () => { + await writeState({ + model: { code: saved }, + variant: { [savedKey]: savedVar }, + }) + const get = spyOn(Agent, "get").mockImplementation(async (name: string) => { + if (name === "code") { + return { + name: "code", + mode: "primary", + permission: [], + options: {}, + model: config, + variant: configVar, + } as any + } + return undefined as any + }) + const modelSpy = spyOn(Provider, "getModel").mockImplementation(async (providerID: string, modelID: string) => { + if (providerID === saved.providerID && modelID === saved.modelID) return savedFull + if (providerID === config.providerID && modelID === config.modelID) return configFull + throw new Error(`unexpected model lookup ${providerID}/${modelID}`) + }) + using _ = { + [Symbol.dispose]() { + get.mockRestore() + modelSpy.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(saved) + expect(user.info.variant).toBe(savedVar) + })) + 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" } } }) @@ -448,13 +536,19 @@ describe("plan follow-up", () => { permission: [], options: {}, model: config, + variant: configVar, } as any } return undefined as any }) + const modelSpy = spyOn(Provider, "getModel").mockImplementation(async (providerID: string, modelID: string) => { + if (providerID === "missing" && modelID === "ghost") throw new Error("missing model") + return configFull + }) using _ = { [Symbol.dispose]() { get.mockRestore() + modelSpy.mockRestore() }, } const seeded = await seed({ text: "1. Build\n2. Test" }) @@ -479,6 +573,7 @@ describe("plan follow-up", () => { if (!user || user.info.role !== "user") return expect(user.info.agent).toBe("code") expect(user.info.model).toEqual(config) + expect(user.info.variant).toBe(configVar) })) test("ask - falls back to planning model when no saved or configured code model exists", () => @@ -492,7 +587,7 @@ describe("plan follow-up", () => { get.mockRestore() }, } - const seeded = await seed({ text: "1. Build\n2. Test" }) + const seeded = await seed({ text: "1. Build\n2. Test", variant: planVar }) const pending = PlanFollowup.ask({ sessionID: seeded.sessionID, messages: seeded.messages, @@ -514,6 +609,7 @@ describe("plan follow-up", () => { if (!user || user.info.role !== "user") return expect(user.info.agent).toBe("code") expect(user.info.model).toEqual(model) + expect(user.info.variant).toBe(planVar) })) test("ask - new session omits handover section when LLM returns empty", () =>