diff --git a/src/plugin/event.model-fallback.test.ts b/src/plugin/event.model-fallback.test.ts index ee7f526a1..a711ef08b 100644 --- a/src/plugin/event.model-fallback.test.ts +++ b/src/plugin/event.model-fallback.test.ts @@ -240,7 +240,7 @@ describe("createEventHandler - model fallback", () => { await Promise.all([messageUpdated, sessionError]) //#then - expect(pendingFallbackArms).toBe(2) + expect(pendingFallbackArms).toBe(1) expect(promptAsyncCalls).toEqual([sessionID]) expect(abortCalls).toEqual([sessionID]) }) @@ -300,7 +300,7 @@ describe("createEventHandler - model fallback", () => { }) //#then - expect(pendingFallbackArms).toBe(2) + expect(pendingFallbackArms).toBe(1) expect(promptAsyncCalls).toEqual([sessionID]) expect(abortCalls).toEqual([sessionID]) }) @@ -517,7 +517,7 @@ describe("createEventHandler - model fallback", () => { expect(promptCalls).toEqual([sessionID]) }) - test("does not re-arm fallback when a duplicate error reports the same failed model after fallback was applied", async () => { + test("does not leave stale pending fallback when a providerless duplicate arrives after fallback was applied", async () => { //#given const sessionID = "ses_model_fallback_duplicate_surface" setMainSession(sessionID) @@ -582,14 +582,12 @@ describe("createEventHandler - model fallback", () => { output, ) - //#when - same failed model arrives again through another OpenCode event surface + //#when - same failed model arrives again without provider metadata after fallback was applied await handler({ event: { type: "session.error", properties: { sessionID, - providerID: "anthropic", - modelID: "claude-opus-4-7-thinking", error: { name: "UnknownError", data: { @@ -603,9 +601,21 @@ describe("createEventHandler - model fallback", () => { }, }) + const staleOutput: ChatMessageOutput = { message: {}, parts: [] } + await chatMessageHandler( + { + sessionID, + agent: "sisyphus", + model: { providerID: "opencode-go", modelID: "kimi-k2.6" }, + }, + staleOutput, + ) + //#then expect(abortCalls).toEqual([sessionID]) expect(promptCalls).toEqual([sessionID]) + expect(modelFallback.hasPendingModelFallback(sessionID)).toBe(false) + expect(staleOutput.message["model"]).toBeUndefined() }) test("does not trigger model-fallback from session.status when runtime_fallback is enabled", async () => { diff --git a/src/plugin/event.ts b/src/plugin/event.ts index 5feaca791..976c732ab 100644 --- a/src/plugin/event.ts +++ b/src/plugin/event.ts @@ -65,6 +65,13 @@ type FallbackContinuationDedupeState = { providerlessModelKeys: Set; }; +type FallbackContinuationContext = { + agentName?: string; + providerID?: string; + dedupeProviderID?: string; + modelID?: string; +}; + function isRecord(value: unknown): value is Record { return typeof value === "object" && value !== null; } @@ -372,12 +379,7 @@ export function createEventHandler(args: { return true; }; - const getFallbackContinuationKeys = (fallbackContext?: { - agentName?: string; - providerID?: string; - dedupeProviderID?: string; - modelID?: string; - }): FallbackContinuationDedupeKeys => { + const getFallbackContinuationKeys = (fallbackContext?: FallbackContinuationContext): FallbackContinuationDedupeKeys => { const agentKey = fallbackContext?.agentName ? getAgentConfigKey(fallbackContext.agentName).trim().toLowerCase() : ""; @@ -424,21 +426,16 @@ export function createEventHandler(args: { return state.providerModelKeys.has(keys.providerModelKey) || state.providerlessModelKeys.has(keys.modelKey); }; - const autoContinueAfterFallback = async ( + const shouldSkipFallbackContinuation = ( sessionID: string, source: string, - fallbackContext?: { - agentName?: string; - providerID?: string; - dedupeProviderID?: string; - modelID?: string; - }, - ): Promise => { + fallbackContext?: FallbackContinuationContext, + ): boolean => { const fallbackKeys = getFallbackContinuationKeys(fallbackContext); if (modelFallbackContinuationsInFlight.has(sessionID)) { log("[event] model-fallback continuation skipped because one is already in flight", { sessionID, source }); - return; + return true; } const lastDispatchedKeys = lastDispatchedModelFallbackContinuationKeys.get(sessionID); @@ -447,6 +444,20 @@ export function createEventHandler(args: { sessionID, source, }); + return true; + } + + return false; + }; + + const autoContinueAfterFallback = async ( + sessionID: string, + source: string, + fallbackContext?: FallbackContinuationContext, + ): Promise => { + const fallbackKeys = getFallbackContinuationKeys(fallbackContext); + + if (shouldSkipFallbackContinuation(sessionID, source, fallbackContext)) { return; } @@ -750,24 +761,26 @@ export function createEventHandler(args: { const currentProvider = resolveFallbackProviderID(sessionID, providerHint); const rawModel = (info?.modelID as string | undefined) ?? "claude-opus-4-7"; const currentModel = normalizeFallbackModelID(rawModel); - applyUserConfiguredFallbackChain(modelFallback, sessionID, agentName, currentProvider, args.pluginConfig); + const fallbackContext = { + agentName, + providerID: currentProvider, + dedupeProviderID: providerHint, + modelID: currentModel, + }; + const shouldAutoContinue = shouldAutoRetrySession(sessionID) && + !hooks.stopContinuationGuard?.isStopped(sessionID); - const setFallback = modelFallback - ? setPendingModelFallback(modelFallback, sessionID, agentName, currentProvider, currentModel) - : false; + if (!shouldAutoContinue || !shouldSkipFallbackContinuation(sessionID, "message.updated", fallbackContext)) { + applyUserConfiguredFallbackChain(modelFallback, sessionID, agentName, currentProvider, args.pluginConfig); - if ( - setFallback && - shouldAutoRetrySession(sessionID) && - !hooks.stopContinuationGuard?.isStopped(sessionID) - ) { - lastHandledModelErrorMessageID.set(sessionID, assistantMessageID); - await autoContinueAfterFallback(sessionID, "message.updated", { - agentName, - providerID: currentProvider, - dedupeProviderID: providerHint, - modelID: currentModel, - }); + const setFallback = modelFallback + ? setPendingModelFallback(modelFallback, sessionID, agentName, currentProvider, currentModel) + : false; + + if (setFallback && shouldAutoContinue) { + lastHandledModelErrorMessageID.set(sessionID, assistantMessageID); + await autoContinueAfterFallback(sessionID, "message.updated", fallbackContext); + } } } } @@ -821,23 +834,25 @@ export function createEventHandler(args: { const currentProvider = resolveFallbackProviderID(sessionID, parsed.providerID); let currentModel = parsed.modelID ?? lastKnown?.modelID ?? "claude-opus-4-7"; currentModel = normalizeFallbackModelID(currentModel); - applyUserConfiguredFallbackChain(modelFallback, sessionID, agentName, currentProvider, args.pluginConfig); + const fallbackContext = { + agentName, + providerID: currentProvider, + dedupeProviderID: parsed.providerID, + modelID: currentModel, + }; + const shouldAutoContinue = shouldAutoRetrySession(sessionID) && + !hooks.stopContinuationGuard?.isStopped(sessionID); - const setFallback = modelFallback - ? setPendingModelFallback(modelFallback, sessionID, agentName, currentProvider, currentModel) - : false; + if (!shouldAutoContinue || !shouldSkipFallbackContinuation(sessionID, "session.status", fallbackContext)) { + applyUserConfiguredFallbackChain(modelFallback, sessionID, agentName, currentProvider, args.pluginConfig); - if ( - setFallback && - shouldAutoRetrySession(sessionID) && - !hooks.stopContinuationGuard?.isStopped(sessionID) - ) { - await autoContinueAfterFallback(sessionID, "session.status", { - agentName, - providerID: currentProvider, - dedupeProviderID: parsed.providerID, - modelID: currentModel, - }); + const setFallback = modelFallback + ? setPendingModelFallback(modelFallback, sessionID, agentName, currentProvider, currentModel) + : false; + + if (setFallback && shouldAutoContinue) { + await autoContinueAfterFallback(sessionID, "session.status", fallbackContext); + } } } } @@ -912,23 +927,25 @@ export function createEventHandler(args: { const currentProvider = resolveFallbackProviderID(sessionID, providerHint); let currentModel = (props?.modelID as string) || parsed.modelID || "claude-opus-4-7"; currentModel = normalizeFallbackModelID(currentModel); - applyUserConfiguredFallbackChain(modelFallback, sessionID, agentName, currentProvider, args.pluginConfig); + const fallbackContext = { + agentName, + providerID: currentProvider, + dedupeProviderID: providerHint, + modelID: currentModel, + }; + const shouldAutoContinue = shouldAutoRetrySession(sessionID) && + !hooks.stopContinuationGuard?.isStopped(sessionID); - const setFallback = modelFallback - ? setPendingModelFallback(modelFallback, sessionID, agentName, currentProvider, currentModel) - : false; + if (!shouldAutoContinue || !shouldSkipFallbackContinuation(sessionID, "session.error", fallbackContext)) { + applyUserConfiguredFallbackChain(modelFallback, sessionID, agentName, currentProvider, args.pluginConfig); - if ( - setFallback && - shouldAutoRetrySession(sessionID) && - !hooks.stopContinuationGuard?.isStopped(sessionID) - ) { - await autoContinueAfterFallback(sessionID, "session.error", { - agentName, - providerID: currentProvider, - dedupeProviderID: providerHint, - modelID: currentModel, - }); + const setFallback = modelFallback + ? setPendingModelFallback(modelFallback, sessionID, agentName, currentProvider, currentModel) + : false; + + if (setFallback && shouldAutoContinue) { + await autoContinueAfterFallback(sessionID, "session.error", fallbackContext); + } } } }