diff --git a/src/hooks/runtime-fallback/index.test.ts b/src/hooks/runtime-fallback/index.test.ts index 8e338bdbb..bb0dbb66b 100644 --- a/src/hooks/runtime-fallback/index.test.ts +++ b/src/hooks/runtime-fallback/index.test.ts @@ -2917,6 +2917,80 @@ describe("runtime-fallback", () => { expect(skipLog).toBeDefined() }) + test("#given a dispatched fallback retry #when stale original assistant error arrives before duplicate session.error #then only one assistant retry prompt is sent", async () => { + const promptCalls: Array = [] + + const hook = createRuntimeFallbackHook( + createMockPluginInput({ + session: { + messages: async () => ({ + data: [{ info: { role: "user" }, parts: [{ type: "text", text: "hello" }] }], + }), + promptAsync: async (args: unknown) => { + promptCalls.push(args) + return {} + }, + }, + }), + { + config: createMockConfig({ notify_on_fallback: false }), + pluginConfig: { + git_master: { + commit_footer: true, + include_co_authored_by: true, + git_env_prefix: "GIT_MASTER=1", + }, + categories: { + test: { + fallback_models: ["provider-a/model-a", "provider-b/model-b"], + }, + }, + }, + } + ) + const sessionID = "test-race-stale-message-update-before-error" + SessionCategoryRegistry.register(sessionID, "test") + + await hook.event({ + event: { + type: "session.created", + properties: { info: { id: sessionID, model: "google/gemini-2.5-pro" } }, + }, + }) + + await hook.event({ + event: { + type: "session.error", + properties: { sessionID, model: "google/gemini-2.5-pro", error: { statusCode: 429, message: "Rate limit" } }, + }, + }) + await hook.event({ + event: { + type: "message.updated", + properties: { + info: { + sessionID, + role: "assistant", + model: "google/gemini-2.5-pro", + error: { statusCode: 429, message: "Rate limit" }, + }, + }, + }, + }) + await hook.event({ + event: { + type: "session.error", + properties: { sessionID, model: "google/gemini-2.5-pro", error: { statusCode: 429, message: "Rate limit" } }, + }, + }) + + expect(promptCalls).toHaveLength(1) + const fallbackLogs = logCalls.filter((call) => call.msg.includes("Preparing fallback")) + expect(fallbackLogs).toHaveLength(1) + const skipLog = logCalls.find((call) => call.msg.includes("session.error skipped - awaiting fallback result")) + expect(skipLog).toBeDefined() + }) + test("session.stop aborts when sessionAwaitingFallbackResult is set", async () => { const abortCalls: Array<{ path?: { id?: string } }> = [] diff --git a/src/hooks/runtime-fallback/message-update-handler.ts b/src/hooks/runtime-fallback/message-update-handler.ts index 00b5f6cc9..88c2a69c4 100644 --- a/src/hooks/runtime-fallback/message-update-handler.ts +++ b/src/hooks/runtime-fallback/message-update-handler.ts @@ -58,7 +58,7 @@ export function createMessageUpdateHandler(deps: HookDeps, helpers: AutoRetryHel sessionAwaitingFallbackResult.delete(sessionID) sessionStatusRetryKeys.delete(sessionID) helpers.clearSessionFallbackTimeout(sessionID) - const state = sessionStates.get(sessionID) + let state = sessionStates.get(sessionID) if (state?.pendingFallbackModel) { state.pendingFallbackModel = undefined } @@ -67,7 +67,25 @@ export function createMessageUpdateHandler(deps: HookDeps, helpers: AutoRetryHel } if (sessionID && role === "assistant" && error) { - const wasAwaitingFallbackResult = sessionAwaitingFallbackResult.delete(sessionID) + let state = sessionStates.get(sessionID) + const pendingFallbackModel = state?.pendingFallbackModel + const wasAwaitingFallbackResult = sessionAwaitingFallbackResult.has(sessionID) + if ( + wasAwaitingFallbackResult && + pendingFallbackModel && + !retrySignal && + model !== pendingFallbackModel + ) { + log(`[${HOOK_NAME}] message.updated fallback skipped - awaiting fallback result`, { + sessionID, + pendingFallbackModel, + model, + }) + return + } + if (wasAwaitingFallbackResult) { + sessionAwaitingFallbackResult.delete(sessionID) + } if (sessionRetryInFlight.has(sessionID) && !retrySignal) { log(`[${HOOK_NAME}] message.updated fallback skipped (retry in flight)`, { sessionID }) return @@ -108,7 +126,6 @@ export function createMessageUpdateHandler(deps: HookDeps, helpers: AutoRetryHel return } - let state = sessionStates.get(sessionID) const agent = info?.agent as string | undefined const resolvedAgent = await helpers.resolveAgentForSessionFromContext(sessionID, agent) const fallbackModels = getFallbackModelsForSession(sessionID, resolvedAgent, pluginConfig)