fix(fallback): guard duplicate prompt injections

This commit is contained in:
YeonGyu-Kim
2026-05-14 01:03:19 +09:00
parent ea55c385bb
commit 5ffbe0e24e
7 changed files with 257 additions and 15 deletions
@@ -39,11 +39,14 @@ import { _resetForTesting as resetSessionState, updateSessionAgent } from "../..
import { runAggressiveTruncationStrategy } from "./aggressive-truncation-strategy"
type FakeClient = {
session: { promptAsync: (input: PromptAsyncCall) => Promise<unknown> }
session: {
promptAsync: (input: PromptAsyncCall) => Promise<unknown>
status?: () => Promise<unknown>
}
tui: { showToast: (input: unknown) => Promise<unknown> }
}
function createRecordingClient(): { client: FakeClient; calls: PromptAsyncCall[] } {
function createRecordingClient(status?: () => Promise<unknown>): { client: FakeClient; calls: PromptAsyncCall[] } {
const calls: PromptAsyncCall[] = []
const client: FakeClient = {
session: {
@@ -51,6 +54,7 @@ function createRecordingClient(): { client: FakeClient; calls: PromptAsyncCall[]
calls.push(input)
return undefined
},
...(status ? { status } : {}),
},
tui: {
showToast: async () => undefined,
@@ -173,4 +177,27 @@ describe("runAggressiveTruncationStrategy - pins agent/model/variant on recovere
expect(calls[0].body.variant).toBeUndefined()
expect(calls[0].body.auto).toBe(true)
})
test("does not send the delayed auto prompt when the session becomes active before recovery fires", async () => {
// given
const sessionID = "session-truncation-active"
const { client, calls } = createRecordingClient(async () => ({
[sessionID]: { type: "busy" },
}))
// when
await runAggressiveTruncationStrategy({
sessionID,
autoCompactState: createAutoCompactState(),
client: client as never,
directory: "/tmp/test-truncation",
truncateAttempt: 0,
currentTokens: 250_000,
maxTokens: 200_000,
})
await flushDeferredPrompt()
// then
expect(calls).toHaveLength(0)
})
})
@@ -17,6 +17,7 @@ import {
findNearestMessageWithFields,
findNearestMessageWithFieldsFromSDK,
} from "../../features/hook-message-injector"
import { isSessionActive } from "../shared/session-idle-settle"
export async function runAggressiveTruncationStrategy(params: {
sessionID: string
@@ -73,6 +74,13 @@ export async function runAggressiveTruncationStrategy(params: {
clearSessionState(params.autoCompactState, params.sessionID)
setTimeout(async () => {
try {
if (await isSessionActive(params.client, params.sessionID)) {
log("[auto-compact] skipped delayed auto prompt because session became active", {
sessionID: params.sessionID,
})
return
}
const sdkMessage = await findNearestMessageWithFieldsFromSDK(params.client, params.sessionID)
const previousMessage = sdkMessage ?? (() => {
const messageDir = getMessageDir(params.sessionID)
@@ -98,7 +106,12 @@ export async function runAggressiveTruncationStrategy(params: {
} as never,
query: { directory: params.directory },
})
} catch {}
} catch (error) {
log("[auto-compact] delayed auto prompt failed", {
sessionID: params.sessionID,
error: String(error),
})
}
}, 500)
return { handled: true, nextTruncateAttempt }
@@ -12,6 +12,19 @@ type ModelFallbackStateLike = {
pending: boolean
}
function canonicalizeModelIDForDuplicateCheck(modelID: string): string {
return modelID.toLowerCase().replace(/\./g, "-")
}
function isSameFailedModel(
state: ModelFallbackStateLike,
providerID: string,
modelID: string,
): boolean {
return state.providerID.toLowerCase() === providerID.toLowerCase()
&& canonicalizeModelIDForDuplicateCheck(state.modelID) === canonicalizeModelIDForDuplicateCheck(modelID)
}
export type ModelFallbackStateController = {
lastToastKey: Map<string, string>
setSessionFallbackChain: (sessionID: string, fallbackChain: FallbackEntry[] | undefined) => void
@@ -84,6 +97,11 @@ export function createModelFallbackStateController(input: {
return false
}
if (existing.attemptCount > 0 && isSameFailedModel(existing, currentProviderID, currentModelID)) {
log(`[model-fallback] Ignoring duplicate fallback arm for already handled model in session: ${sessionID}`)
return false
}
existing.providerID = currentProviderID
existing.modelID = currentModelID
existing.pending = true
@@ -12,6 +12,21 @@ import { dispatchFallbackRetry } from "./fallback-retry-dispatcher"
import { createSessionStatusHandler } from "./session-status-handler"
import { resolveMessageEventSessionID, resolveSessionEventID } from "../../shared/event-session-id"
function resolveEventModel(props: Record<string, unknown> | undefined): string | undefined {
const model = props?.model
if (typeof model === "string") {
return model
}
const providerID = props?.providerID
const modelID = props?.modelID
if (typeof providerID === "string" && typeof modelID === "string") {
return `${providerID}/${modelID}`
}
return undefined
}
export function createEventHandler(deps: HookDeps, helpers: AutoRetryHelpers) {
const { config, pluginConfig, sessionStates, sessionLastAccess, sessionRetryInFlight, sessionAwaitingFallbackResult, sessionFallbackTimeouts, sessionStatusRetryKeys } = deps
const sessionStatusHandler = createSessionStatusHandler(deps, helpers, sessionStatusRetryKeys)
@@ -137,6 +152,19 @@ export function createEventHandler(deps: HookDeps, helpers: AutoRetryHelpers) {
return
}
if (sessionAwaitingFallbackResult.has(sessionID)) {
const pendingFallbackModel = sessionStates.get(sessionID)?.pendingFallbackModel
const eventModel = resolveEventModel(props)
if (!pendingFallbackModel || eventModel !== pendingFallbackModel) {
log(`[${HOOK_NAME}] session.error skipped - awaiting fallback result`, {
sessionID,
pendingFallbackModel,
eventModel,
})
return
}
}
sessionAwaitingFallbackResult.delete(sessionID)
helpers.clearSessionFallbackTimeout(sessionID)
+66 -1
View File
@@ -2650,7 +2650,7 @@ describe("runtime-fallback", () => {
await hook.event({
event: {
type: "session.error",
properties: { sessionID, error: { statusCode: 429, message: "Rate limit again" } },
properties: { sessionID, model: "provider-a/model-a", error: { statusCode: 429, message: "Rate limit again" } },
},
})
@@ -2659,6 +2659,71 @@ describe("runtime-fallback", () => {
expect(fallbackLogs.length).toBeGreaterThanOrEqual(2)
})
test("session.error is skipped while waiting for the dispatched fallback result", async () => {
const promptCalls: Array<unknown> = []
//#given
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-awaiting-fallback-result"
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, error: { statusCode: 429, message: "Rate limit" } },
},
})
//#when - duplicate stale error fires after promptAsync resolved but before fallback output is visible
await hook.event({
event: {
type: "session.error",
properties: { sessionID, error: { statusCode: 429, message: "Rate limit" } },
},
})
//#then
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 } }> = []