fix(cli): preserve selected variant as well

This commit is contained in:
Josh Holmer
2026-03-19 10:44:56 -04:00
parent 1fc0e644dd
commit 040284034a
2 changed files with 148 additions and 23 deletions
+49 -20
View File
@@ -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<MessageV2.User, "model" | "variant">) {
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<string, MessageV2.User["model"]>
variant?: Record<string, string | undefined>
},
)
.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"
@@ -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<void>) {
await using tmp = await tmpdir({ git: true })
@@ -52,6 +58,7 @@ async function withInstance(fn: () => Promise<void>) {
async function seed(input: {
text: string
variant?: string
tools?: Array<{ tool: string; input: Record<string, unknown>; 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<string, { providerID: string; modelID: string }> }) {
async function writeState(input: {
model?: Record<string, { providerID: string; modelID: string }>
variant?: Record<string, string | undefined>
}) {
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", () =>