Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions .changeset/preserve-model-variants.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
---
"kilo-code": patch
---

Preserve the selected reasoning effort when switching to a model that supports the same or nearest available variant.
9 changes: 8 additions & 1 deletion packages/kilo-vscode/tests/unit/mode-model.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -13,8 +13,15 @@ describe("modelPatch", () => {
})
})

it("clears stale variant when next model does not support it", () => {
it("keeps the nearest supported effort when the exact variant is unavailable", () => {
expect(modelPatch("kilo", "anthropic/claude-sonnet-4-6", ["low", "medium"], "high")).toEqual({
model: "kilo/anthropic/claude-sonnet-4-6",
variant: "medium",
})
})

it("clears an unknown variant when next model does not support it", () => {
expect(modelPatch("kilo", "anthropic/claude-sonnet-4-6", ["low", "medium"], "thinking")).toEqual({
model: "kilo/anthropic/claude-sonnet-4-6",
variant: null,
})
Expand Down
22 changes: 22 additions & 0 deletions packages/kilo-vscode/tests/unit/session-variant-store.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ import {
cycleVariant,
getAgentVariant,
getVariant,
preserveVariant,
sessionVariantKeys,
sessionVariants,
transferVariants,
Expand Down Expand Up @@ -121,3 +122,24 @@ describe("cycleVariant", () => {
expect(cycleVariant("low", [])).toBeUndefined()
})
})

describe("preserveVariant", () => {
it("keeps an exact variant", () => {
expect(preserveVariant("high", ["low", "high"])).toBe("high")
expect(preserveVariant("thinking", ["instant", "thinking"])).toBe("thinking")
expect(preserveVariant("default", ["default", "thinking"])).toBe("default")
})

it("falls back to the nearest supported effort", () => {
expect(preserveVariant("max", ["high", "xhigh"])).toBe("xhigh")
expect(preserveVariant("high", ["low", "medium"])).toBe("medium")
expect(preserveVariant("max", ["none", "low"])).toBe("low")
})

it("does not cross binary or custom variant families", () => {
expect(preserveVariant("thinking", ["low", "high"])).toBeUndefined()
expect(preserveVariant("instant", ["low", "high"])).toBeUndefined()
expect(preserveVariant("turbo", ["low", "high"])).toBeUndefined()
expect(preserveVariant("high", ["instant", "thinking"])).toBeUndefined()
})
})
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@ import { useServer } from "../src/context/server"
import { useSession } from "../src/context/session"
import { useProvider } from "../src/context/provider"
import { useConfig } from "../src/context/config"
import { cycleVariant } from "../src/context/session-variant-store"
import { cycleVariant, preserveVariant } from "../src/context/session-variant-store"
import { ModelSelectorBase } from "../src/components/shared/ModelSelector"
import { ModeSwitcherBase } from "../src/components/shared/ModeSwitcher"
import { SpeechToTextButton } from "../src/components/speech-to-text/SpeechToTextButton"
Expand Down Expand Up @@ -268,7 +268,7 @@ export const NewWorktreeDialog: Component<{
return
}
const stored = variant()
if (!stored || !list.includes(stored)) setVariant(list[0])
if (!stored || !list.includes(stored)) setVariant(preserveVariant(stored, list) ?? list[0])
})

createEffect(() => {
Expand Down Expand Up @@ -871,7 +871,12 @@ export const NewWorktreeDialog: Component<{
<ModelSelectorBase
value={model()}
onSelect={(pid, mid) => {
if (pid && mid) setModel({ providerID: pid, modelID: mid })
if (!pid || !mid) return
const current = effectiveVariant()
const next = { providerID: pid, modelID: mid }
const list = Object.keys(provider.findModel(next)?.variants ?? {})
setModel(next)
setVariant(preserveVariant(current, list))
}}
onPick={restorePrompt}
onCancel={restorePrompt}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ import { DEFAULT_SPEECH_TO_TEXT_MODEL } from "../../../../src/speech-to-text/mod
import { hasSpeechToTextAccess, selectedSpeechToTextModel } from "../speech-to-text/availability"
import { speechToTextModelOptions } from "../speech-to-text/model-selector"
import { AUTOCOMPLETE_SELECTOR_MODELS, getAutocompleteSelection } from "./autocomplete-model-selector"
import { preserveVariant } from "../../context/session-variant-store"

const ModelsTab: Component = () => {
const { config, settings, updateConfig, updateSetting } = useConfig()
Expand Down Expand Up @@ -64,9 +65,12 @@ const ModelsTab: Component = () => {
return
}
const value = `${providerID}/${modelID}`
const list = Object.keys(provider.findModel({ providerID, modelID })?.variants ?? {})
const next = preserveVariant(subagentVariant(), list)
updateConfig({
subagent_model: value,
...(config().subagent_model === value ? {} : { subagent_variant: null }),
...(next ? { subagent_variant_overrides: { ...config().subagent_variant_overrides, [value]: next } } : {}),
})
}

Expand All @@ -87,7 +91,17 @@ const ModelsTab: Component = () => {
updateConfig({ agent: { [agentName]: { model: null } } })
return
}
updateConfig({ agent: { [agentName]: { model: `${providerID}/${modelID}` } } })
const current = config().agent?.[agentName]?.variant ?? undefined
const list = Object.keys(provider.findModel({ providerID, modelID })?.variants ?? {})
const next = preserveVariant(current, list)
updateConfig({
agent: {
[agentName]: {
model: `${providerID}/${modelID}`,
...(current && !list.includes(current) ? { variant: next ?? null } : {}),
},
},
})
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ import { ModelSelectorBase } from "../../shared/ModelSelector"
import { ThinkingSelectorBase } from "../../shared/ThinkingSelector"
import { parseModelString } from "../../../../../src/shared/provider-model"
import type { CommandConfig } from "../../../types/messages"
import { preserveVariant } from "../../../context/session-variant-store"

const WorkflowsTab: Component = () => {
const language = useLanguage()
Expand Down Expand Up @@ -48,9 +49,10 @@ const WorkflowsTab: Component = () => {
const selectModel = (name: string, providerID: string, modelID: string) => {
const list = Object.keys(provider.findModel({ providerID, modelID })?.variants ?? {})
const current = variant(config().command?.[name] ?? {}, name)
const next = preserveVariant(current, list)
update(name, {
model: providerID && modelID ? `${providerID}/${modelID}` : null,
...(current && !list.includes(current) ? { variant: null } : {}),
...(current && !list.includes(current) ? { variant: next ?? null } : {}),
})
}

Expand Down
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import type { AgentConfig } from "../../types/messages"
import { preserveVariant } from "../../context/session-variant-store"

export function modelPatch(
providerID: string,
Expand All @@ -12,6 +13,6 @@ export function modelPatch(

return {
model: `${providerID}/${modelID}`,
...(current && !variants.includes(current) ? { variant: null } : {}),
...(current && !variants.includes(current) ? { variant: preserveVariant(current, variants) ?? null } : {}),
}
}
Original file line number Diff line number Diff line change
@@ -1,5 +1,26 @@
import type { ModelSelection } from "../types/messages"

const effort = ["none", "minimal", "low", "medium", "high", "xhigh", "max"]

/** Keep the selected effort when possible, falling back to the nearest known effort. */
export function preserveVariant(current: string | undefined, variants: string[]) {
if (!current || variants.length === 0) return undefined
if (variants.includes(current)) return current

const rank = effort.indexOf(current)
if (rank === -1) return undefined

return variants
.map((value, index) => ({
value,
index,
rank: effort.indexOf(value),
distance: Math.abs(effort.indexOf(value) - rank),
}))
.filter((item) => effort.includes(item.value))
.sort((a, b) => a.distance - b.distance || b.rank - a.rank || a.index - b.index)[0]?.value
}

export function legacyVariantKey(sel: ModelSelection) {
return `${sel.providerID}/${sel.modelID}`
}
Expand All @@ -21,7 +42,7 @@ export function getVariant(
const key = variantKey(sel, agent, session)
const fallback = session ? store[variantKey(sel, agent)] : undefined
const stored = store[key] ?? fallback ?? store[legacyVariantKey(sel)]
return stored && variants.includes(stored) ? stored : variants[0]
return preserveVariant(stored, variants) ?? variants[0]
}

export function getAgentVariant(
Expand Down
28 changes: 26 additions & 2 deletions packages/kilo-vscode/webview-ui/src/context/session.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,14 @@
import { PartStash } from "./part-stash"
import { mergeParts, sameParts } from "./session-parts"
import { state as todoState } from "./todo-revert"
import { getAgentVariant, getVariant, sessionVariantKeys, transferVariants, variantKey } from "./session-variant-store"
import {
getAgentVariant,
getVariant,
preserveVariant,
sessionVariantKeys,
transferVariants,
variantKey,
} from "./session-variant-store"
import { KILO_AUTO, KILO_PROVIDER_ID, parseModelString } from "../../../src/shared/provider-model"
import { reviewMetadata, type ReviewMessageData } from "../../../src/shared/review-comments"
import { visibleMessages as filterVisibleMessages } from "./session-queue"
Expand Down Expand Up @@ -663,9 +670,22 @@
})
}

function carryVariant(selection: ModelSelection, current: string | undefined, agent: string, sessionID?: string) {
const value = preserveVariant(current, Object.keys(provider.findModel(selection)?.variants ?? {}))
if (!value) return
const key = variantKey(selection, agent, sessionID)
setStore("variantSelections", key, value)
if (!sessionID) vscode.postMessage({ type: "persistVariant", key, value })
}

function selectModel(providerID: string, modelID: string, sessionID?: string) {
const sid = sessionID ?? currentSessionID()
applyModel(agentForScope(sid), { providerID, modelID }, sid)
const agent = agentForScope(sid)
const current = selected(sid)
const value = current ? currentVariant(sid) : undefined
const selection = { providerID, modelID }
applyModel(agent, selection, sid)
carryVariant(selection, value, agent, sid)
if (sid) {
hideErrors(sid)
}
Expand Down Expand Up @@ -3033,8 +3053,12 @@
// agent may not yet be assigned (sendInitialMessage calls setSessionModel
// before setSessionAgent), so the write would land on defaultAgent() and
// corrupt the default mode's model for later sessions.
const agent = store.agentSelections[sessionID] ?? defaultAgent()
const current = selected(sessionID)
const value = current ? currentVariant(sessionID) : undefined
const model = { providerID, modelID }
setStore("sessionOverrides", sessionID, model)
carryVariant(model, value, agent, sessionID)
},
setSessionAgent: (sessionID: string, name: string) => {
setStore("agentSelections", sessionID, name)
Expand Down Expand Up @@ -3074,7 +3098,7 @@
dismissSuggestion,
createSession,
clearCurrentSession,
loadSessions,

Check failure on line 3101 in packages/kilo-vscode/webview-ui/src/context/session.tsx

View workflow job for this annotation

GitHub Actions / unit tests

File has too many lines (3124). Maximum allowed is 3100
loadOlderMessages,
selectSession,
deleteSession,
Expand Down
4 changes: 0 additions & 4 deletions packages/opencode/src/cli/cmd/run/footer.ts
Original file line number Diff line number Diff line change
Expand Up @@ -834,11 +834,7 @@ export class RunFooter implements FooterApi {
return
}

const previous = this.currentModel()
this.setCurrentModel(model)
if (!previous || previous.providerID !== model.providerID || previous.modelID !== model.modelID) {
this.setCurrentVariant(undefined)
}
void Promise.resolve()
.then(() => this.options.onModelSelect?.(model))
.then((result) => {
Expand Down
17 changes: 14 additions & 3 deletions packages/opencode/src/cli/cmd/run/runtime.ts
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,8 @@ import { resolveModelInfo, resolveRunTuiConfig, resolveSessionInfo } from "./run
import { createRuntimeLifecycle } from "./runtime.lifecycle"
import { trace } from "./trace"
import { cycleVariant, formatModelLabel, resolveSavedVariant, resolveVariant, saveVariant } from "./variant.shared"
// kilocode_change - preserve compatible variants when switching models
import { resolvePreservedVariant } from "@/kilocode/cli/cmd/run/variant" // kilocode_change
import type { LocalReplayAnchor, LocalReplayRow, RunInput, RunPrompt, RunProvider, StreamCommit } from "./types"

/** @internal Exported for testing */
Expand Down Expand Up @@ -294,17 +296,22 @@ async function runInteractiveRuntime(input: RunRuntimeInput, deps: RunRuntimeDep
return
}

// kilocode_change start - preserve the active effort across model switches
const previous = state.activeVariant
state.model = model
state.activeVariant = undefined
state.variants = variantsFor(state.providers, model)
const switching = resolveSavedVariant(model).then((saved) => {
const current = state.model
if (!current || current.providerID !== model.providerID || current.modelID !== model.modelID) {
return
}

state.activeVariant = resolveVariant(ctx.variant, undefined, saved, state.variants)
// kilocode_change - prefer the active effort over a model-specific saved preference
state.activeVariant =
resolvePreservedVariant(ctx.variant, previous, state.variants) ??
resolveVariant(ctx.variant, undefined, saved, state.variants)
})
// kilocode_change end
state.switching = switching
await switching
if (state.switching === switching) {
Expand Down Expand Up @@ -440,10 +447,14 @@ async function runInteractiveRuntime(input: RunRuntimeInput, deps: RunRuntimeDep
state.variants = variantsFor(state.providers, state.model)
state.limits = info.limits

const next = resolveVariant(ctx.variant, session.variant, savedVariant, state.variants)
// kilocode_change start - preserve the active effort when the model catalog arrives asynchronously
const next =
resolvePreservedVariant(ctx.variant, state.activeVariant, state.variants) ??
resolveVariant(ctx.variant, session.variant, savedVariant, state.variants)
if (next !== state.activeVariant) {
state.activeVariant = next
}
// kilocode_change end

if (footer.isClosed) {
return
Expand Down
29 changes: 29 additions & 0 deletions packages/opencode/src/kilocode/cli/cmd/run/variant.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
const effort = ["none", "minimal", "low", "medium", "high", "xhigh", "max"]

/** Keep an explicit CLI variant verbatim; only infer a fallback for automatic selections. */
export function resolvePreservedVariant(
input: string | undefined,
current: string | undefined,
variants: string[],
): string | undefined {
return input ?? preserveVariant(current, variants)
}

/** Keep the selected effort when possible, falling back to the nearest known effort. */
export function preserveVariant(current: string | undefined, variants: string[]): string | undefined {
if (!current || variants.length === 0) return undefined
if (variants.includes(current)) return current

const rank = effort.indexOf(current)
if (rank === -1) return undefined

return variants
.map((value, index) => ({
value,
index,
rank: effort.indexOf(value),
distance: Math.abs(effort.indexOf(value) - rank),
}))
.filter((item) => effort.includes(item.value))
.sort((a, b) => a.distance - b.distance || b.rank - a.rank || a.index - b.index)[0]?.value
}
28 changes: 28 additions & 0 deletions packages/opencode/test/kilocode/cli/run/variant.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
import { describe, expect, test } from "bun:test"
import { preserveVariant, resolvePreservedVariant } from "@/kilocode/cli/cmd/run/variant"

describe("Kilo CLI variant preservation", () => {
test("keeps exact variants across supported families", () => {
expect(preserveVariant("high", ["low", "high"])).toBe("high")
expect(preserveVariant("thinking", ["instant", "thinking"])).toBe("thinking")
expect(preserveVariant("default", ["default", "thinking"])).toBe("default")
})

test("keeps an explicit CLI variant verbatim", () => {
expect(resolvePreservedVariant("max", "high", ["low", "medium", "high"])).toBe("max")
expect(resolvePreservedVariant("thinking", "high", ["low", "medium", "high"])).toBe("thinking")
})

test("falls back to the nearest supported reasoning effort", () => {
expect(preserveVariant("max", ["high", "xhigh"])).toBe("xhigh")
expect(preserveVariant("high", ["low", "medium"])).toBe("medium")
expect(preserveVariant("max", ["none", "low"])).toBe("low")
})

test("does not cross binary or custom variant families", () => {
expect(preserveVariant("thinking", ["low", "high"])).toBeUndefined()
expect(preserveVariant("instant", ["low", "high"])).toBeUndefined()
expect(preserveVariant("turbo", ["low", "high"])).toBeUndefined()
expect(preserveVariant("high", ["instant", "thinking"])).toBeUndefined()
})
})
Loading