Compare commits

...

1 Commits

Author SHA1 Message Date
Aiden Cline d841feb3e3 fix(tui): follow agent model selection 2026-08-18 14:29:55 +00:00
3 changed files with 106 additions and 83 deletions
+1 -7
View File
@@ -1232,7 +1232,6 @@ export function Prompt(props: PromptProps) {
} else if (slashHead && isCommand) {
move.startSubmit()
const model = { providerID: selection.providerID, id: selection.modelID, variant }
const cancelCommit = local.model.trackSessionCommit(sessionID, model)
void client.api.session
.command({
@@ -1247,7 +1246,6 @@ export function Prompt(props: PromptProps) {
delivery,
})
.catch((error) => {
cancelCommit()
toast.show({ title: "Failed to run command", message: errorMessage(error), variant: "error" })
})
} else if (isSkill) {
@@ -1271,11 +1269,7 @@ export function Prompt(props: PromptProps) {
(session.model.variant ?? "default") !== (variant ?? "default")
) {
const model = { providerID: selection.providerID, id: selection.modelID, variant }
const cancelCommit = local.model.trackSessionCommit(sessionID, model)
await client.api.session.switchModel({ sessionID, model }).catch((error) => {
cancelCommit()
throw error
})
await client.api.session.switchModel({ sessionID, model })
}
if (session?.revert) {
const error = await client.api.session.revert.commit({ sessionID }).then(
+60 -75
View File
@@ -45,6 +45,39 @@ export function recentModels(model: ModelPreferenceModel, recent: ModelPreferenc
.map((item) => ({ providerID: item.providerID, modelID: item.modelID }))
}
type ModelSelection = ModelPreferenceModel & { variant?: string }
type AgentSelection = {
id: string
model?: { providerID: string; id: string; variant?: string }
}
type SessionSelection = {
agent?: string
model?: { providerID: string; id: string; variant?: string }
}
export function resolveAgentModelSelection(input: {
selected?: ModelSelection
agent?: AgentSelection
session?: SessionSelection
available: (model: ModelPreferenceModel) => boolean
}) {
const model = (value: SessionSelection["model"]): ModelSelection | undefined =>
value && {
providerID: value.providerID,
modelID: value.id,
variant: normalizeModelVariant(value.variant),
}
const candidates = [
input.selected,
input.session?.agent === input.agent?.id ? model(input.session?.model) : undefined,
model(input.agent?.model),
model(input.session?.model),
]
return candidates.find((item): item is ModelSelection => !!item && input.available(item))
}
export const { use: useLocal, provider: LocalProvider } = createSimpleContext({
name: "Local",
init: () => {
@@ -67,13 +100,6 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({
return !!models()?.some((item) => item.providerID === model.providerID && item.id === model.modelID)
}
function getFirstValidModel(...modelFns: (() => ModelPreferenceModel | undefined)[]) {
for (const modelFn of modelFns) {
const model = modelFn()
if (model && isModelValid(model)) return model
}
}
function createAgent() {
const agents = createMemo(() =>
(data.location.agent.list(location.ref) ?? []).filter((agent) => agent.mode !== "subagent" && !agent.hidden),
@@ -132,7 +158,6 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({
const agent = createAgent()
function createModel() {
type ModelSelection = ModelPreferenceModel & { variant?: string }
const [preferences, setPreferences] = createStore<ModelPreference & { ready: boolean }>({
ready: false,
recent: [],
@@ -141,16 +166,13 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({
})
const [selectionState, setSelectionState] = createStore<{
newSessionModelByLocationAgent: Record<string, ModelPreferenceModel | undefined>
draftBySession: Record<string, ModelSelection | undefined>
modelBySessionAgent: Record<string, Record<string, ModelSelection | undefined> | undefined>
}>({
newSessionModelByLocationAgent: {},
draftBySession: {},
modelBySessionAgent: {},
})
const repository = createModelPreferenceRepository(path.join(paths.state, "model.json"))
const pendingSelectionCommits = new Map<string, string>()
const selectionKey = (value: ModelSelection) =>
`${modelPreferenceKey(value)}:${normalizeModelVariant(value.variant) ?? "default"}`
const saveState = {
pending: false,
}
@@ -208,20 +230,19 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({
}
})
const newSessionModel = createMemo(() => {
const newSessionSelection = createMemo<ModelSelection | undefined>(() => {
const a = agent.current()
return getFirstValidModel(
() => a && selectionState.newSessionModelByLocationAgent[locationAgentKey(a.id)],
() => a?.model && { providerID: a.model.providerID, modelID: a.model.id },
fallbackModel,
)
const selected = a && selectionState.newSessionModelByLocationAgent[locationAgentKey(a.id)]
const resolved = resolveAgentModelSelection({ selected, agent: a, available: isModelValid }) ?? fallbackModel()
if (!resolved) return
if (selected || !a?.model || resolved.providerID !== a.model.providerID || resolved.modelID !== a.model.id)
return { ...resolved, variant: normalizeModelVariant(preferences.variant[modelPreferenceKey(resolved)]) }
return resolved
})
const currentSelection = createMemo<ModelSelection | undefined>(() => {
if (route.data.type === "session") return sessionSelection(route.data.sessionID)
const model = newSessionModel()
if (!model) return
return { ...model, variant: normalizeModelVariant(preferences.variant[modelPreferenceKey(model)]) }
return newSessionSelection()
})
const currentModel = createMemo(() => {
@@ -235,27 +256,23 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({
return `${JSON.stringify([ref.directory, ref.workspaceID])}:${agentID}`
}
function durableSelection(sessionID: string): ModelSelection | undefined {
const model = data.session.get(sessionID)?.model
if (!model) return
return {
providerID: model.providerID,
modelID: model.id,
variant: normalizeModelVariant(model.variant),
}
}
function sessionSelection(sessionID: string) {
return selectionState.draftBySession[sessionID] ?? durableSelection(sessionID)
const current = agent.current()
return resolveAgentModelSelection({
selected: current && selectionState.modelBySessionAgent[sessionID]?.[current.id],
agent: current,
session: data.session.get(sessionID),
available: isModelValid,
})
}
function setSessionDraft(sessionID: string, selection: ModelSelection) {
const durable = durableSelection(sessionID)
setSelectionState(
"draftBySession",
sessionID,
durable && selectionKey(durable) === selectionKey(selection) ? undefined : selection,
)
function setSessionSelection(sessionID: string, selection: ModelSelection) {
const current = agent.current()
if (!current) return
setSelectionState("modelBySessionAgent", sessionID, {
...selectionState.modelBySessionAgent[sessionID],
[current.id]: selection,
})
}
function selectModel(model: ModelPreferenceModel) {
@@ -269,7 +286,7 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({
)
const info = models()?.find((item) => item.providerID === model.providerID && item.id === model.modelID)
const variant = preferred && info?.variants?.some((item) => item.id === preferred) ? preferred : undefined
setSessionDraft(sessionID, { ...model, variant })
setSessionSelection(sessionID, { ...model, variant })
return true
}
const current = agent.current()
@@ -278,27 +295,9 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({
return true
}
onCleanup(
event.on("session.model.selected", (evt) => {
const expected = pendingSelectionCommits.get(evt.data.sessionID)
if (!expected) return
const committed = selectionKey({
providerID: evt.data.model.providerID,
modelID: evt.data.model.id,
variant: evt.data.model.variant,
})
if (committed !== expected) return
pendingSelectionCommits.delete(evt.data.sessionID)
const draft = selectionState.draftBySession[evt.data.sessionID]
if (draft && selectionKey(draft) === committed)
setSelectionState("draftBySession", evt.data.sessionID, undefined)
}),
)
onCleanup(
event.on("session.deleted", (evt) => {
pendingSelectionCommits.delete(evt.data.sessionID)
setSelectionState("draftBySession", evt.data.sessionID, undefined)
setSelectionState("modelBySessionAgent", evt.data.sessionID, undefined)
}),
)
@@ -308,20 +307,6 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({
available(model = currentModel()) {
return model ? isModelValid(model) : false
},
trackSessionCommit(
sessionID: string,
value: {
providerID: string
id: string
variant?: string
},
) {
const committed = selectionKey({ providerID: value.providerID, modelID: value.id, variant: value.variant })
pendingSelectionCommits.set(sessionID, committed)
return () => {
if (pendingSelectionCommits.get(sessionID) === committed) pendingSelectionCommits.delete(sessionID)
}
},
get ready() {
return preferences.ready
},
@@ -436,7 +421,7 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({
const m = currentSelection()
if (!m) return
if (route.data.type === "session") {
setSessionDraft(route.data.sessionID, { ...m, variant: normalizeModelVariant(value) })
setSessionSelection(route.data.sessionID, { ...m, variant: normalizeModelVariant(value) })
}
setPreferences("variant", modelPreferenceKey(m), normalizeModelVariant(value))
savePreferences()
+45 -1
View File
@@ -1,5 +1,5 @@
import { expect, test } from "bun:test"
import { parseModel, recentModels } from "../../src/context/local"
import { parseModel, recentModels, resolveAgentModelSelection } from "../../src/context/local"
test("parses model IDs containing slashes", () => {
expect(parseModel("provider/family/model")).toEqual({
@@ -20,3 +20,47 @@ test("moves a model to the front, deduplicates, and limits recents", () => {
...recent.slice(6, 10),
])
})
test("uses the configured model when switching agents", () => {
expect(
resolveAgentModelSelection({
agent: { id: "build", model: { providerID: "provider", id: "build-model", variant: "max" } },
session: {
agent: "plan",
model: { providerID: "provider", id: "plan-model", variant: "high" },
},
available: () => true,
}),
).toEqual({ providerID: "provider", modelID: "build-model", variant: "max" })
})
test("keeps a manual model selection for each session agent", () => {
expect(
resolveAgentModelSelection({
selected: { providerID: "provider", modelID: "manual-model", variant: "high" },
agent: { id: "build", model: { providerID: "provider", id: "build-model", variant: "max" } },
session: { agent: "plan", model: { providerID: "provider", id: "plan-model" } },
available: () => true,
}),
).toEqual({ providerID: "provider", modelID: "manual-model", variant: "high" })
})
test("keeps the durable model while the active agent is unchanged", () => {
expect(
resolveAgentModelSelection({
agent: { id: "plan", model: { providerID: "provider", id: "configured-model", variant: "max" } },
session: { agent: "plan", model: { providerID: "provider", id: "manual-model", variant: "high" } },
available: () => true,
}),
).toEqual({ providerID: "provider", modelID: "manual-model", variant: "high" })
})
test("keeps the session model when the next agent has no configured model", () => {
expect(
resolveAgentModelSelection({
agent: { id: "review" },
session: { agent: "plan", model: { providerID: "provider", id: "plan-model", variant: "high" } },
available: () => true,
}),
).toEqual({ providerID: "provider", modelID: "plan-model", variant: "high" })
})