diff --git a/src/plugin/event.model-fallback.test.ts b/src/plugin/event.model-fallback.test.ts index 15ed40131..f75b99682 100644 --- a/src/plugin/event.model-fallback.test.ts +++ b/src/plugin/event.model-fallback.test.ts @@ -301,6 +301,68 @@ describe("createEventHandler - model fallback", () => { expect(abortCalls).toEqual([sessionID]) }) + test("does not collapse fallback continuations for different providers with the same model id", async () => { + //#given + const sessionID = "ses_model_fallback_same_model_different_provider" + setMainSession(sessionID) + let pendingFallbackArms = 0 + const modelFallback = unsafeTestValue({ + setSessionFallbackChain: () => {}, + setPendingModelFallback: () => { + pendingFallbackArms += 1 + return true + }, + }) + const { handler, abortCalls, promptAsyncCalls } = createHandler({ + hooks: { modelFallback }, + promptAsync: async () => ({}), + }) + + const assistantError = { + name: "APIError", + data: { + message: + "Bad Gateway: {\"error\":{\"message\":\"unknown provider for model claude-opus-4-7-thinking\"}}", + isRetryable: true, + }, + } + + await handler({ + event: { + type: "message.updated", + properties: { + info: { + id: "msg_err_same_model_provider_1", + sessionID, + role: "assistant", + error: assistantError, + modelID: "claude-opus-4-7-thinking", + providerID: "anthropic", + agent: "Sisyphus - Ultraworker", + }, + }, + }, + }) + + //#when - a distinct provider reports the same normalized model id before idle cleanup + await handler({ + event: { + type: "session.error", + properties: { + sessionID, + providerID: "quotio", + modelID: "claude-opus-4-7-thinking", + error: assistantError, + }, + }, + }) + + //#then + expect(pendingFallbackArms).toBe(2) + expect(promptAsyncCalls).toEqual([sessionID, sessionID]) + expect(abortCalls).toEqual([sessionID, sessionID]) + }) + test("triggers retry prompt on session.status retry events and applies fallback", async () => { //#given const sessionID = "ses_status_retry_fallback" diff --git a/src/plugin/event.ts b/src/plugin/event.ts index eeefe0e78..5feaca791 100644 --- a/src/plugin/event.ts +++ b/src/plugin/event.ts @@ -54,6 +54,17 @@ type FirstMessageVariantGate = { clear: (sessionID: string) => void; }; +type FallbackContinuationDedupeKeys = { + modelKey?: string; + providerModelKey?: string; +}; + +type FallbackContinuationDedupeState = { + modelKeys: Set; + providerModelKeys: Set; + providerlessModelKeys: Set; +}; + function isRecord(value: unknown): value is Record { return typeof value === "object" && value !== null; } @@ -220,7 +231,7 @@ export function createEventHandler(args: { const lastHandledRetryStatusKey = new Map(); const lastKnownModelBySession = new Map(); const modelFallbackContinuationsInFlight = new Set(); - const lastDispatchedModelFallbackContinuationKeys = new Map>(); + const lastDispatchedModelFallbackContinuationKeys = new Map(); const resolveFallbackProviderID = (sessionID: string, providerHint?: string): string => { const normalizedProviderHint = providerHint?.trim(); @@ -364,23 +375,53 @@ export function createEventHandler(args: { const getFallbackContinuationKeys = (fallbackContext?: { agentName?: string; providerID?: string; + dedupeProviderID?: string; modelID?: string; - }): string[] => { + }): FallbackContinuationDedupeKeys => { const agentKey = fallbackContext?.agentName ? getAgentConfigKey(fallbackContext.agentName).trim().toLowerCase() : ""; - const providerID = fallbackContext?.providerID?.trim().toLowerCase() ?? ""; + const providerID = fallbackContext?.dedupeProviderID?.trim().toLowerCase() ?? ""; const modelID = fallbackContext?.modelID?.trim().toLowerCase() ?? ""; if (!agentKey || !modelID) { - return []; + return {}; } - const keys = [`${agentKey}:${modelID}`]; - if (providerID) { - keys.push(`${agentKey}:${providerID}:${modelID}`); + return { + modelKey: `${agentKey}:${modelID}`, + ...(providerID ? { providerModelKey: `${agentKey}:${providerID}:${modelID}` } : {}), + }; + }; + + const getFallbackContinuationDedupeState = (sessionID: string): FallbackContinuationDedupeState => { + const existingState = lastDispatchedModelFallbackContinuationKeys.get(sessionID); + if (existingState) { + return existingState; } - return keys; + + const state = { + modelKeys: new Set(), + providerModelKeys: new Set(), + providerlessModelKeys: new Set(), + }; + lastDispatchedModelFallbackContinuationKeys.set(sessionID, state); + return state; + }; + + const wasFallbackContinuationAlreadyDispatched = ( + state: FallbackContinuationDedupeState | undefined, + keys: FallbackContinuationDedupeKeys, + ): boolean => { + if (!state || !keys.modelKey) { + return false; + } + + if (!keys.providerModelKey) { + return state.modelKeys.has(keys.modelKey); + } + + return state.providerModelKeys.has(keys.providerModelKey) || state.providerlessModelKeys.has(keys.modelKey); }; const autoContinueAfterFallback = async ( @@ -389,6 +430,7 @@ export function createEventHandler(args: { fallbackContext?: { agentName?: string; providerID?: string; + dedupeProviderID?: string; modelID?: string; }, ): Promise => { @@ -400,7 +442,7 @@ export function createEventHandler(args: { } const lastDispatchedKeys = lastDispatchedModelFallbackContinuationKeys.get(sessionID); - if (lastDispatchedKeys && fallbackKeys.some((fallbackKey) => lastDispatchedKeys.has(fallbackKey))) { + if (wasFallbackContinuationAlreadyDispatched(lastDispatchedKeys, fallbackKeys)) { log("[event] model-fallback continuation skipped because matching fallback was already dispatched", { sessionID, source, @@ -456,12 +498,14 @@ export function createEventHandler(args: { log("[event] model-fallback prompt failed", { sessionID, source, error }); }); } finally { - if (dispatched && fallbackKeys.length > 0) { - const dispatchedKeys = lastDispatchedModelFallbackContinuationKeys.get(sessionID) ?? new Set(); - for (const fallbackKey of fallbackKeys) { - dispatchedKeys.add(fallbackKey); + if (dispatched && fallbackKeys.modelKey) { + const dispatchedKeys = getFallbackContinuationDedupeState(sessionID); + dispatchedKeys.modelKeys.add(fallbackKeys.modelKey); + if (fallbackKeys.providerModelKey) { + dispatchedKeys.providerModelKeys.add(fallbackKeys.providerModelKey); + } else { + dispatchedKeys.providerlessModelKeys.add(fallbackKeys.modelKey); } - lastDispatchedModelFallbackContinuationKeys.set(sessionID, dispatchedKeys); } modelFallbackContinuationsInFlight.delete(sessionID); } @@ -702,10 +746,8 @@ export function createEventHandler(args: { } if (agentName) { - const currentProvider = resolveFallbackProviderID( - sessionID, - info?.providerID as string | undefined, - ); + const providerHint = info?.providerID as string | undefined; + 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); @@ -723,6 +765,7 @@ export function createEventHandler(args: { await autoContinueAfterFallback(sessionID, "message.updated", { agentName, providerID: currentProvider, + dedupeProviderID: providerHint, modelID: currentModel, }); } @@ -792,6 +835,7 @@ export function createEventHandler(args: { await autoContinueAfterFallback(sessionID, "session.status", { agentName, providerID: currentProvider, + dedupeProviderID: parsed.providerID, modelID: currentModel, }); } @@ -864,10 +908,8 @@ export function createEventHandler(args: { if (agentName) { const parsed = extractProviderModelFromErrorMessage(errorMessage); - const currentProvider = resolveFallbackProviderID( - sessionID, - (props?.providerID as string | undefined) || parsed.providerID, - ); + const providerHint = (props?.providerID as string | undefined) || parsed.providerID; + 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); @@ -884,6 +926,7 @@ export function createEventHandler(args: { await autoContinueAfterFallback(sessionID, "session.error", { agentName, providerID: currentProvider, + dedupeProviderID: providerHint, modelID: currentModel, }); }