diff --git a/src/config/schema/hooks.ts b/src/config/schema/hooks.ts index e3bad81a8..fea9c6371 100644 --- a/src/config/schema/hooks.ts +++ b/src/config/schema/hooks.ts @@ -25,6 +25,7 @@ export const HookNameSchema = z.enum([ "interactive-bash-session", "thinking-block-validator", + "tool-pair-validator", "ralph-loop", "category-skill-reminder", diff --git a/src/hooks/index.ts b/src/hooks/index.ts index 3966a5bea..051cbd12a 100644 --- a/src/hooks/index.ts +++ b/src/hooks/index.ts @@ -26,6 +26,7 @@ export { createNonInteractiveEnvHook } from "./non-interactive-env"; export { createInteractiveBashSessionHook } from "./interactive-bash-session"; export { createThinkingBlockValidatorHook } from "./thinking-block-validator"; +export { createToolPairValidatorHook } from "./tool-pair-validator"; export { createCategorySkillReminderHook } from "./category-skill-reminder"; export { createRalphLoopHook, type RalphLoopHook } from "./ralph-loop"; export { createNoSisyphusGptHook } from "./no-sisyphus-gpt"; diff --git a/src/hooks/tool-pair-validator/hook.test.ts b/src/hooks/tool-pair-validator/hook.test.ts new file mode 100644 index 000000000..6b18f15f0 --- /dev/null +++ b/src/hooks/tool-pair-validator/hook.test.ts @@ -0,0 +1,156 @@ +declare const describe: (name: string, fn: () => void) => void +declare const it: (name: string, fn: () => void | Promise) => void +declare const expect: (value: T) => { + toEqual(expected: unknown): void + toHaveLength(expected: number): void +} + +import { createToolPairValidatorHook } from "./hook" + +const TOOL_RESULT_PLACEHOLDER = "Tool output unavailable (context compacted)" + +type TestPart = { + type: string + id?: string + callID?: string + tool_use_id?: string + content?: string + text?: string +} + +type TestMessage = { + info: { role: "assistant" | "user" } + parts: TestPart[] +} + +async function runTransform(messages: TestMessage[]): Promise { + const hook = createToolPairValidatorHook() + const transform = hook["experimental.chat.messages.transform"] + + if (!transform) { + throw new Error("missing tool pair validator transform") + } + + await transform({}, { messages: messages as never }) +} + +describe("createToolPairValidatorHook", () => { + it("leaves matching tool pairs unchanged", async () => { + //#given + const messages = [ + { info: { role: "assistant" }, parts: [{ type: "tool", callID: "call_1" }] }, + { info: { role: "user" }, parts: [{ type: "tool_result", tool_use_id: "call_1", content: "done" }] }, + ] satisfies TestMessage[] + + //#when + await runTransform(messages) + + //#then + expect(messages).toEqual([ + { info: { role: "assistant" }, parts: [{ type: "tool", callID: "call_1" }] }, + { info: { role: "user" }, parts: [{ type: "tool_result", tool_use_id: "call_1", content: "done" }] }, + ]) + }) + + it("injects a missing tool_result into the next user message", async () => { + //#given + const messages = [ + { info: { role: "assistant" }, parts: [{ type: "tool_use", id: "toolu_1" }] }, + { info: { role: "user" }, parts: [{ type: "text", text: "continue" }] }, + ] satisfies TestMessage[] + + //#when + await runTransform(messages) + + //#then + expect(messages[1]?.parts).toEqual([ + { type: "tool_result", tool_use_id: "toolu_1", content: TOOL_RESULT_PLACEHOLDER }, + { type: "text", text: "continue" }, + ]) + }) + + it("injects a synthetic user message when the next user message is missing", async () => { + //#given + const messages = [ + { + info: { role: "assistant" }, + parts: [ + { type: "tool_use", id: "toolu_1" }, + { type: "text", text: "working" }, + { type: "tool_use", id: "toolu_2" }, + ], + }, + ] satisfies TestMessage[] + + //#when + await runTransform(messages) + + //#then + expect(messages).toEqual([ + { + info: { role: "assistant" }, + parts: [ + { type: "tool_use", id: "toolu_1" }, + { type: "text", text: "working" }, + { type: "tool_use", id: "toolu_2" }, + ], + }, + { + info: { role: "user" }, + parts: [ + { type: "tool_result", tool_use_id: "toolu_1", content: TOOL_RESULT_PLACEHOLDER }, + { type: "tool_result", tool_use_id: "toolu_2", content: TOOL_RESULT_PLACEHOLDER }, + ], + }, + ]) + }) + + it("injects a synthetic user message before a non-user next message", async () => { + //#given + const messages = [ + { info: { role: "assistant" }, parts: [{ type: "tool_use", id: "toolu_1" }] }, + { info: { role: "assistant" }, parts: [{ type: "text", text: "follow-up" }] }, + ] satisfies TestMessage[] + + //#when + await runTransform(messages) + + //#then + expect(messages).toHaveLength(3) + expect(messages).toEqual([ + { info: { role: "assistant" }, parts: [{ type: "tool_use", id: "toolu_1" }] }, + { + info: { role: "user" }, + parts: [{ type: "tool_result", tool_use_id: "toolu_1", content: TOOL_RESULT_PLACEHOLDER }], + }, + { info: { role: "assistant" }, parts: [{ type: "text", text: "follow-up" }] }, + ]) + }) + + it("injects only the missing tool_results for partial matches", async () => { + //#given + const messages = [ + { + info: { role: "assistant" }, + parts: [{ type: "tool_use", id: "toolu_1" }, { type: "tool", callID: "call_2" }], + }, + { + info: { role: "user" }, + parts: [ + { type: "tool_result", tool_use_id: "toolu_1", content: "done" }, + { type: "text", text: "continue" }, + ], + }, + ] satisfies TestMessage[] + + //#when + await runTransform(messages) + + //#then + expect(messages[1]?.parts).toEqual([ + { type: "tool_result", tool_use_id: "toolu_1", content: "done" }, + { type: "tool_result", tool_use_id: "call_2", content: TOOL_RESULT_PLACEHOLDER }, + { type: "text", text: "continue" }, + ]) + }) +}) diff --git a/src/hooks/tool-pair-validator/hook.ts b/src/hooks/tool-pair-validator/hook.ts new file mode 100644 index 000000000..89a76e701 --- /dev/null +++ b/src/hooks/tool-pair-validator/hook.ts @@ -0,0 +1,184 @@ +import type { Message, Part } from "@opencode-ai/sdk" + +import { log } from "../../shared/logger" + +const TOOL_RESULT_PLACEHOLDER = "Tool output unavailable (context compacted)" + +type ToolUsePart = { + type: "tool_use" + id: string + [key: string]: unknown +} + +type ToolResultPart = { + type: "tool_result" + tool_use_id: string + content: string + [key: string]: unknown +} + +type TransformPart = Part | ToolUsePart | ToolResultPart + +type TransformMessageInfo = Message | { + role: "user" + sessionID?: string +} + +interface MessageWithParts { + info: TransformMessageInfo + parts: TransformPart[] +} + +type MessagesTransformHook = { + "experimental.chat.messages.transform"?: ( + input: Record, + output: { messages: MessageWithParts[] } + ) => Promise +} + +function getToolUseID(part: TransformPart): string | null { + const candidate = part as { type?: unknown; id?: unknown; callID?: unknown } + + if (candidate.type === "tool_use" && typeof candidate.id === "string" && candidate.id.length > 0) { + return candidate.id + } + + if (candidate.type === "tool" && typeof candidate.callID === "string" && candidate.callID.length > 0) { + return candidate.callID + } + + return null +} + +function getToolResultID(part: TransformPart): string | null { + const candidate = part as { type?: unknown; tool_use_id?: unknown } + + if (candidate.type === "tool_result" && typeof candidate.tool_use_id === "string" && candidate.tool_use_id.length > 0) { + return candidate.tool_use_id + } + + return null +} + +function extractUniqueToolUseIDs(parts: TransformPart[]): string[] { + const seen = new Set() + const toolUseIDs: string[] = [] + + for (const part of parts) { + const toolUseID = getToolUseID(part) + if (!toolUseID || seen.has(toolUseID)) { + continue + } + + seen.add(toolUseID) + toolUseIDs.push(toolUseID) + } + + return toolUseIDs +} + +function extractToolResultIDs(parts: TransformPart[]): Set { + const toolResultIDs = new Set() + + for (const part of parts) { + const toolResultID = getToolResultID(part) + if (toolResultID) { + toolResultIDs.add(toolResultID) + } + } + + return toolResultIDs +} + +function createToolResultPart(toolUseID: string): ToolResultPart { + return { + type: "tool_result", + tool_use_id: toolUseID, + content: TOOL_RESULT_PLACEHOLDER, + } +} + +function findToolResultInsertIndex(parts: TransformPart[]): number { + let lastToolResultIndex = -1 + + for (let i = 0; i < parts.length; i++) { + if (getToolResultID(parts[i])) { + lastToolResultIndex = i + } + } + + return lastToolResultIndex === -1 ? 0 : lastToolResultIndex + 1 +} + +function insertMissingToolResults(message: MessageWithParts, missingToolUseIDs: string[]): void { + const toolResultParts = missingToolUseIDs.map((toolUseID) => createToolResultPart(toolUseID)) + const insertIndex = findToolResultInsertIndex(message.parts) + message.parts.splice(insertIndex, 0, ...toolResultParts) +} + +function createSyntheticUserMessage(assistantMessage: MessageWithParts, missingToolUseIDs: string[]): MessageWithParts { + const assistantInfo = assistantMessage.info as { sessionID?: unknown } + const sessionID = typeof assistantInfo.sessionID === "string" ? assistantInfo.sessionID : undefined + + return { + info: { + role: "user", + ...(sessionID ? { sessionID } : {}), + }, + parts: missingToolUseIDs.map((toolUseID) => createToolResultPart(toolUseID)), + } +} + +function getMessageID(message: TransformMessageInfo): string | undefined { + const candidate = message as { id?: unknown } + return typeof candidate.id === "string" ? candidate.id : undefined +} + +function repairMissingToolResults(messages: MessageWithParts[], assistantIndex: number): void { + const assistantMessage = messages[assistantIndex] + const toolUseIDs = extractUniqueToolUseIDs(assistantMessage.parts) + + if (toolUseIDs.length === 0) { + return + } + + const nextMessage = messages[assistantIndex + 1] + + if (nextMessage?.info.role !== "user") { + messages.splice(assistantIndex + 1, 0, createSyntheticUserMessage(assistantMessage, toolUseIDs)) + log("[tool-pair-validator] Repaired missing tool_result blocks", { + assistantMessageID: getMessageID(assistantMessage.info), + syntheticUserMessageInserted: true, + repairedToolUseIDs: toolUseIDs, + }) + return + } + + const existingToolResultIDs = extractToolResultIDs(nextMessage.parts) + const missingToolUseIDs = toolUseIDs.filter((toolUseID) => !existingToolResultIDs.has(toolUseID)) + + if (missingToolUseIDs.length === 0) { + return + } + + insertMissingToolResults(nextMessage, missingToolUseIDs) + log("[tool-pair-validator] Repaired missing tool_result blocks", { + assistantMessageID: getMessageID(assistantMessage.info), + syntheticUserMessageInserted: false, + repairedToolUseIDs: missingToolUseIDs, + }) +} + +export function createToolPairValidatorHook(): MessagesTransformHook { + return { + "experimental.chat.messages.transform": async (_input, output) => { + for (let i = 0; i < output.messages.length; i++) { + if (output.messages[i].info.role !== "assistant") { + continue + } + + repairMissingToolResults(output.messages, i) + } + }, + } +} diff --git a/src/hooks/tool-pair-validator/index.ts b/src/hooks/tool-pair-validator/index.ts new file mode 100644 index 000000000..717bede96 --- /dev/null +++ b/src/hooks/tool-pair-validator/index.ts @@ -0,0 +1 @@ +export { createToolPairValidatorHook } from "./hook" diff --git a/src/plugin/hooks/create-transform-hooks.ts b/src/plugin/hooks/create-transform-hooks.ts index d593efae7..c57a959bb 100644 --- a/src/plugin/hooks/create-transform-hooks.ts +++ b/src/plugin/hooks/create-transform-hooks.ts @@ -5,6 +5,7 @@ import { createClaudeCodeHooksHook, createKeywordDetectorHook, createThinkingBlockValidatorHook, + createToolPairValidatorHook, } from "../../hooks" import { contextCollector, @@ -17,6 +18,7 @@ export type TransformHooks = { keywordDetector: ReturnType | null contextInjectorMessagesTransform: ReturnType thinkingBlockValidator: ReturnType | null + toolPairValidator: ReturnType | null } export function createTransformHooks(args: { @@ -63,10 +65,19 @@ export function createTransformHooks(args: { ) : null + const toolPairValidator = isHookEnabled("tool-pair-validator") + ? safeCreateHook( + "tool-pair-validator", + () => createToolPairValidatorHook(), + { enabled: safeHookEnabled }, + ) + : null + return { claudeCodeHooks, keywordDetector, contextInjectorMessagesTransform, thinkingBlockValidator, + toolPairValidator, } } diff --git a/src/plugin/messages-transform.ts b/src/plugin/messages-transform.ts index 6ea674d8a..cd28b3832 100644 --- a/src/plugin/messages-transform.ts +++ b/src/plugin/messages-transform.ts @@ -20,5 +20,9 @@ export function createMessagesTransformHandler(args: { await args.hooks.thinkingBlockValidator?.[ "experimental.chat.messages.transform" ]?.(input, output) + + await args.hooks.toolPairValidator?.[ + "experimental.chat.messages.transform" + ]?.(input, output) } }