diff --git a/src/hooks/session-recovery/recover-unavailable-tool.ts b/src/hooks/session-recovery/recover-unavailable-tool.ts index 01c6660c9..2b8cf0702 100644 --- a/src/hooks/session-recovery/recover-unavailable-tool.ts +++ b/src/hooks/session-recovery/recover-unavailable-tool.ts @@ -21,6 +21,17 @@ interface PromptWithToolResultInput { body: { parts: ToolResultPart[] } } +type ClientWithPromptAsync = Client & { + session: Client["session"] & { + promptAsync: (input: PromptWithToolResultInput) => Promise + } +} + +function hasPromptAsync(client: Client): client is ClientWithPromptAsync { + const promptAsync = (client.session as { promptAsync?: unknown }).promptAsync + return typeof promptAsync === "function" +} + interface ToolUsePart { type: "tool_use" id: string @@ -104,17 +115,12 @@ export async function recoverUnavailableTool( path: { id: sessionID }, body: { parts: toolResultParts }, } - const promptAsync = client.session.promptAsync as (...args: never[]) => unknown - const promptClient = { - session: { - status: client.session.status, - promptAsync: (input: PromptWithToolResultInput) => ( - Reflect.apply(promptAsync, client.session, [input]) as Promise - ), - }, + if (!hasPromptAsync(client)) { + return false } + const promptResult = await promptAsyncAfterSessionIdle({ - client: promptClient, + client, sessionID, source: "session-recovery-unavailable-tool", input: promptInput, diff --git a/src/shared/prompt-async-route-audit.test.ts b/src/shared/prompt-async-route-audit.test.ts index 0335b7c47..32b00f88b 100644 --- a/src/shared/prompt-async-route-audit.test.ts +++ b/src/shared/prompt-async-route-audit.test.ts @@ -42,6 +42,14 @@ describe("production prompt injection routes", () => { // given const files = await listSourceFiles(SOURCE_ROOT) const offenders: string[] = [] + const rawPromptPatterns = [ + /\bsession\.promptAsync\s*\(/, + /\bsession\.prompt\s*\(/, + /\bReflect\.apply\s*\(\s*\w*promptAsync\b/, + /\bReflect\.apply\s*\(\s*\w*prompt\b/, + /\b(?:const|let|var)\s+\w*promptAsync\w*\s*=\s*[\w.]+\.session\.promptAsync\b/, + /\b(?:const|let|var)\s+\w*prompt\w*\s*=\s*[\w.]+\.session\.prompt\b/, + ] // when for (const filePath of files) { @@ -50,7 +58,7 @@ describe("production prompt injection routes", () => { } const contents = uncommentedLines(await readFile(filePath, "utf8")).join("\n") - if (/\bsession\.promptAsync\s*\(/.test(contents) || /\bsession\.prompt\s*\(/.test(contents)) { + if (rawPromptPatterns.some((pattern) => pattern.test(contents))) { offenders.push(relativeSourcePath(filePath)) } }