Merge pull request #3169 from code-yeongyu/fix/sisyphus-junior-model
fix(model-fallback): respect user model override for sisyphus-junior category sessions (#2941)
This commit is contained in:
@@ -332,6 +332,25 @@ describe("model fallback hook", () => {
|
|||||||
clearPendingModelFallback(sessionID)
|
clearPendingModelFallback(sessionID)
|
||||||
})
|
})
|
||||||
|
|
||||||
|
test("does not fall back to hardcoded agent chain when session explicitly stores no fallback chain [regression #2941]", () => {
|
||||||
|
//#given
|
||||||
|
const sessionID = "ses_model_fallback_explicit_none"
|
||||||
|
clearPendingModelFallback(sessionID)
|
||||||
|
setSessionFallbackChain(sessionID, undefined)
|
||||||
|
|
||||||
|
//#when
|
||||||
|
const set = setPendingModelFallback(
|
||||||
|
sessionID,
|
||||||
|
"Sisyphus - Junior",
|
||||||
|
"anthropic",
|
||||||
|
"claude-sonnet-4-6",
|
||||||
|
)
|
||||||
|
|
||||||
|
//#then
|
||||||
|
expect(set).toBe(false)
|
||||||
|
clearPendingModelFallback(sessionID)
|
||||||
|
})
|
||||||
|
|
||||||
test("shows toast when fallback is applied", async () => {
|
test("shows toast when fallback is applied", async () => {
|
||||||
//#given
|
//#given
|
||||||
const toastCalls: Array<{ title: string; message: string }> = []
|
const toastCalls: Array<{ title: string; message: string }> = []
|
||||||
|
|||||||
@@ -42,8 +42,12 @@ const sessionFallbackChains = new Map<string, FallbackEntry[]>()
|
|||||||
|
|
||||||
export function setSessionFallbackChain(sessionID: string, fallbackChain: FallbackEntry[] | undefined): void {
|
export function setSessionFallbackChain(sessionID: string, fallbackChain: FallbackEntry[] | undefined): void {
|
||||||
if (!sessionID) return
|
if (!sessionID) return
|
||||||
if (!fallbackChain || fallbackChain.length === 0) {
|
if (!fallbackChain) {
|
||||||
sessionFallbackChains.delete(sessionID)
|
sessionFallbackChains.set(sessionID, [])
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if (fallbackChain.length === 0) {
|
||||||
|
sessionFallbackChains.set(sessionID, [])
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
sessionFallbackChains.set(sessionID, fallbackChain)
|
sessionFallbackChains.set(sessionID, fallbackChain)
|
||||||
@@ -65,8 +69,9 @@ export function setPendingModelFallback(
|
|||||||
): boolean {
|
): boolean {
|
||||||
const agentKey = getAgentConfigKey(agentName)
|
const agentKey = getAgentConfigKey(agentName)
|
||||||
const requirements = AGENT_MODEL_REQUIREMENTS[agentKey]
|
const requirements = AGENT_MODEL_REQUIREMENTS[agentKey]
|
||||||
|
const hasSessionFallback = sessionFallbackChains.has(sessionID)
|
||||||
const sessionFallback = sessionFallbackChains.get(sessionID)
|
const sessionFallback = sessionFallbackChains.get(sessionID)
|
||||||
const fallbackChain = sessionFallback && sessionFallback.length > 0
|
const fallbackChain = hasSessionFallback
|
||||||
? sessionFallback
|
? sessionFallback
|
||||||
: requirements?.fallbackChain
|
: requirements?.fallbackChain
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,165 @@
|
|||||||
|
declare const require: (name: string) => any
|
||||||
|
const { afterEach, describe, expect, spyOn, test } = require("bun:test")
|
||||||
|
|
||||||
|
import { createEventHandler } from "./event"
|
||||||
|
import { createChatMessageHandler } from "./chat-message"
|
||||||
|
import { _resetForTesting, setSessionAgent } from "../features/claude-code-session-state"
|
||||||
|
import { clearPendingModelFallback, createModelFallbackHook, setSessionFallbackChain } from "../hooks/model-fallback/hook"
|
||||||
|
import * as connectedProvidersCache from "../shared/connected-providers-cache"
|
||||||
|
|
||||||
|
type EventInput = { event: { type: string; properties?: unknown } }
|
||||||
|
type EventHandlerArgs = Parameters<typeof createEventHandler>[0]
|
||||||
|
type EventHandlerInput = Parameters<ReturnType<typeof createEventHandler>>[0]
|
||||||
|
type ChatMessageHandlerArgs = Parameters<typeof createChatMessageHandler>[0]
|
||||||
|
|
||||||
|
function asEventHandlerInput(input: EventInput): EventHandlerInput {
|
||||||
|
return input as unknown as EventHandlerInput
|
||||||
|
}
|
||||||
|
|
||||||
|
function asEventHandlerContext(ctx: unknown): EventHandlerArgs["ctx"] {
|
||||||
|
return ctx as unknown as EventHandlerArgs["ctx"]
|
||||||
|
}
|
||||||
|
|
||||||
|
function asPluginConfig(config: unknown): EventHandlerArgs["pluginConfig"] {
|
||||||
|
return config as unknown as EventHandlerArgs["pluginConfig"]
|
||||||
|
}
|
||||||
|
|
||||||
|
function asChatMessageHandlerContext(ctx: unknown): ChatMessageHandlerArgs["ctx"] {
|
||||||
|
return ctx as unknown as ChatMessageHandlerArgs["ctx"]
|
||||||
|
}
|
||||||
|
|
||||||
|
function asChatPluginConfig(config: unknown): ChatMessageHandlerArgs["pluginConfig"] {
|
||||||
|
return config as unknown as ChatMessageHandlerArgs["pluginConfig"]
|
||||||
|
}
|
||||||
|
|
||||||
|
function createEventHandlerManagers(): EventHandlerArgs["managers"] {
|
||||||
|
return {
|
||||||
|
tmuxSessionManager: {
|
||||||
|
onSessionCreated: async () => {},
|
||||||
|
onSessionDeleted: async () => {},
|
||||||
|
},
|
||||||
|
skillMcpManager: {
|
||||||
|
disconnectSession: async () => {},
|
||||||
|
},
|
||||||
|
} as unknown as EventHandlerArgs["managers"]
|
||||||
|
}
|
||||||
|
|
||||||
|
function createEventHandlerHooks(modelFallback: ReturnType<typeof createModelFallbackHook>): EventHandlerArgs["hooks"] {
|
||||||
|
return {
|
||||||
|
modelFallback,
|
||||||
|
} as unknown as EventHandlerArgs["hooks"]
|
||||||
|
}
|
||||||
|
|
||||||
|
function createChatMessageHandlerHooks(modelFallback: ReturnType<typeof createModelFallbackHook>): ChatMessageHandlerArgs["hooks"] {
|
||||||
|
return {
|
||||||
|
modelFallback,
|
||||||
|
stopContinuationGuard: null,
|
||||||
|
keywordDetector: null,
|
||||||
|
claudeCodeHooks: null,
|
||||||
|
autoSlashCommand: null,
|
||||||
|
startWork: null,
|
||||||
|
ralphLoop: null,
|
||||||
|
} as unknown as ChatMessageHandlerArgs["hooks"]
|
||||||
|
}
|
||||||
|
|
||||||
|
let readConnectedProvidersCacheSpy: { mockRestore: () => void } | undefined
|
||||||
|
let readProviderModelsCacheSpy: { mockRestore: () => void } | undefined
|
||||||
|
|
||||||
|
afterEach(() => {
|
||||||
|
readConnectedProvidersCacheSpy?.mockRestore()
|
||||||
|
readProviderModelsCacheSpy?.mockRestore()
|
||||||
|
readConnectedProvidersCacheSpy = undefined
|
||||||
|
readProviderModelsCacheSpy = undefined
|
||||||
|
_resetForTesting()
|
||||||
|
})
|
||||||
|
|
||||||
|
describe("createEventHandler - category runtime fallback suppression", () => {
|
||||||
|
test("does not arm retry fallback when category session explicitly stores no fallback chain [regression #2941]", async () => {
|
||||||
|
//#given
|
||||||
|
const sessionID = "ses_category_override_no_fallback"
|
||||||
|
const abortCalls: string[] = []
|
||||||
|
const promptCalls: string[] = []
|
||||||
|
|
||||||
|
readConnectedProvidersCacheSpy = spyOn(connectedProvidersCache, "readConnectedProvidersCache").mockReturnValue(null)
|
||||||
|
readProviderModelsCacheSpy = spyOn(connectedProvidersCache, "readProviderModelsCache").mockReturnValue(null)
|
||||||
|
|
||||||
|
clearPendingModelFallback(sessionID)
|
||||||
|
setSessionAgent(sessionID, "sisyphus-junior")
|
||||||
|
setSessionFallbackChain(sessionID, undefined)
|
||||||
|
|
||||||
|
const modelFallback = createModelFallbackHook()
|
||||||
|
const eventHandler = createEventHandler({
|
||||||
|
ctx: asEventHandlerContext({
|
||||||
|
directory: "/tmp",
|
||||||
|
client: {
|
||||||
|
session: {
|
||||||
|
abort: async ({ path }: { path: { id: string } }) => {
|
||||||
|
abortCalls.push(path.id)
|
||||||
|
return {}
|
||||||
|
},
|
||||||
|
prompt: async ({ path }: { path: { id: string } }) => {
|
||||||
|
promptCalls.push(path.id)
|
||||||
|
return {}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}),
|
||||||
|
pluginConfig: asPluginConfig({}),
|
||||||
|
firstMessageVariantGate: {
|
||||||
|
markSessionCreated: () => {},
|
||||||
|
clear: () => {},
|
||||||
|
},
|
||||||
|
managers: createEventHandlerManagers(),
|
||||||
|
hooks: createEventHandlerHooks(modelFallback),
|
||||||
|
})
|
||||||
|
|
||||||
|
const chatMessageHandler = createChatMessageHandler({
|
||||||
|
ctx: asChatMessageHandlerContext({
|
||||||
|
client: {
|
||||||
|
tui: {
|
||||||
|
showToast: async () => ({}),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}),
|
||||||
|
pluginConfig: asChatPluginConfig({}),
|
||||||
|
firstMessageVariantGate: {
|
||||||
|
shouldOverride: () => false,
|
||||||
|
markApplied: () => {},
|
||||||
|
},
|
||||||
|
hooks: createChatMessageHandlerHooks(modelFallback),
|
||||||
|
})
|
||||||
|
|
||||||
|
//#when
|
||||||
|
await eventHandler(asEventHandlerInput({
|
||||||
|
event: {
|
||||||
|
type: "session.error",
|
||||||
|
properties: {
|
||||||
|
sessionID,
|
||||||
|
error: {
|
||||||
|
name: "APIError",
|
||||||
|
data: {
|
||||||
|
message:
|
||||||
|
"Bad Gateway: {\"error\":{\"message\":\"unknown provider for model claude-sonnet-4-6\"}}",
|
||||||
|
isRetryable: true,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}))
|
||||||
|
|
||||||
|
const output = { message: {}, parts: [] as Array<{ type: string; text?: string }> }
|
||||||
|
await chatMessageHandler(
|
||||||
|
{
|
||||||
|
sessionID,
|
||||||
|
agent: "sisyphus-junior",
|
||||||
|
model: { providerID: "anthropic", modelID: "claude-sonnet-4-6" },
|
||||||
|
},
|
||||||
|
output,
|
||||||
|
)
|
||||||
|
|
||||||
|
//#then
|
||||||
|
expect(abortCalls).toEqual([])
|
||||||
|
expect(promptCalls).toEqual([])
|
||||||
|
expect(output.message["model"]).toBeUndefined()
|
||||||
|
})
|
||||||
|
})
|
||||||
@@ -484,4 +484,32 @@ describe("resolveCategoryExecution", () => {
|
|||||||
})
|
})
|
||||||
expect(result.fallbackChain).toBeUndefined()
|
expect(result.fallbackChain).toBeUndefined()
|
||||||
})
|
})
|
||||||
|
|
||||||
|
test("does not inherit hardcoded fallbackChain when sisyphus-junior model override is set [regression #2941]", async () => {
|
||||||
|
//#given
|
||||||
|
const args = {
|
||||||
|
category: "quick",
|
||||||
|
prompt: "test prompt",
|
||||||
|
description: "Test task",
|
||||||
|
run_in_background: false,
|
||||||
|
load_skills: [],
|
||||||
|
blockedBy: undefined,
|
||||||
|
enableSkillTools: false,
|
||||||
|
}
|
||||||
|
const executorCtx = createMockExecutorContext()
|
||||||
|
executorCtx.sisyphusJuniorModel = "anthropic/claude-sonnet-4-6"
|
||||||
|
|
||||||
|
//#when
|
||||||
|
const result = await resolveCategoryExecution(args, executorCtx, undefined, "anthropic/claude-sonnet-4-6")
|
||||||
|
|
||||||
|
//#then
|
||||||
|
expect(result.error).toBeUndefined()
|
||||||
|
expect(result.actualModel).toBe("anthropic/claude-sonnet-4-6")
|
||||||
|
expect(result.categoryModel).toEqual({
|
||||||
|
providerID: "anthropic",
|
||||||
|
modelID: "claude-sonnet-4-6",
|
||||||
|
variant: undefined,
|
||||||
|
})
|
||||||
|
expect(result.fallbackChain).toBeUndefined()
|
||||||
|
})
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -275,6 +275,6 @@ Available categories: ${categoryNames.join(", ")}`,
|
|||||||
actualModel,
|
actualModel,
|
||||||
isUnstableAgent,
|
isUnstableAgent,
|
||||||
// Don't use hardcoded fallback chain when resolution was skipped (cold cache)
|
// Don't use hardcoded fallback chain when resolution was skipped (cold cache)
|
||||||
fallbackChain: configuredFallbackChain ?? ((isModelResolutionSkipped || explicitCategoryModel) ? undefined : requirement?.fallbackChain),
|
fallbackChain: configuredFallbackChain ?? ((isModelResolutionSkipped || explicitCategoryModel || overrideModel) ? undefined : requirement?.fallbackChain),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2963,6 +2963,7 @@ describe("sisyphus-task", () => {
|
|||||||
// then - sisyphus-junior override model should be used, not category default
|
// then - sisyphus-junior override model should be used, not category default
|
||||||
expect(launchInput.model.providerID).toBe("anthropic")
|
expect(launchInput.model.providerID).toBe("anthropic")
|
||||||
expect(launchInput.model.modelID).toBe("claude-sonnet-4-6")
|
expect(launchInput.model.modelID).toBe("claude-sonnet-4-6")
|
||||||
|
expect(launchInput.fallbackChain).toBeUndefined()
|
||||||
})
|
})
|
||||||
|
|
||||||
test("sisyphus-junior model override works with user-defined category (#1295)", async () => {
|
test("sisyphus-junior model override works with user-defined category (#1295)", async () => {
|
||||||
|
|||||||
Reference in New Issue
Block a user