From 5c7a7dcb49b2e3d0a93928435281e2d7d7ec34a9 Mon Sep 17 00:00:00 2001 From: Brace Sproul Date: Thu, 24 Jul 2025 16:23:02 -0700 Subject: [PATCH] fix: Dont diagnose err after just diagnosing err (#527) * fix: Dont diagnose err after just diagnosing err * cr --- .../src/__tests__/take-action.test.ts | 99 ++++++++++++++++--- apps/open-swe/src/utils/tool-message-error.ts | 59 +++++++++-- 2 files changed, 137 insertions(+), 21 deletions(-) diff --git a/apps/open-swe/src/__tests__/take-action.test.ts b/apps/open-swe/src/__tests__/take-action.test.ts index 902875f2..e254c92b 100644 --- a/apps/open-swe/src/__tests__/take-action.test.ts +++ b/apps/open-swe/src/__tests__/take-action.test.ts @@ -157,23 +157,98 @@ describe("Error diagnosis logic", () => { expect(shouldDiagnoseError(messages)).toBe(false); }); - test("should ignore diagnostic tool messages", () => { + test("should ignore diagnostic tool messages when calculating error rates", () => { const messages = [ - new AIMessage({ content: "AI message 1" }), // AI message 1 - createToolMessage("1", "tool1", "error"), - createToolMessage("2", "tool2", "error", true), // Diagnostic (ignored) + new AIMessage({ content: "AI message 0" }), // AI message 0 (NOT part of last 3) + createToolMessage("0", "tool0", "error"), + createToolMessage("1", "diagnose_error", "success", true), // Old diagnostic (ignored and outside last 3) - new AIMessage({ content: "AI message 2" }), // AI message 2 - createToolMessage("3", "tool3", "error"), - createToolMessage("4", "tool4", "error"), + new AIMessage({ content: "AI message 1" }), // AI message 1 (part of last 3) + createToolMessage("2", "tool1", "error"), + createToolMessage("3", "tool2", "error"), - new AIMessage({ content: "AI message 3" }), // AI message 3 - createToolMessage("5", "tool5", "error"), - createToolMessage("6", "tool6", "error"), - createToolMessage("7", "tool7", "error", true), // Diagnostic (ignored) + new AIMessage({ content: "AI message 2" }), // AI message 2 (part of last 3) + createToolMessage("4", "tool3", "error"), + createToolMessage("5", "tool4", "error"), + + new AIMessage({ content: "AI message 3" }), // AI message 3 (part of last 3) + createToolMessage("6", "tool5", "error"), + createToolMessage("7", "tool6", "error"), ]; - expect(shouldDiagnoseError(messages)).toBe(true); // All 3 groups have 100% error rate + expect(shouldDiagnoseError(messages)).toBe(true); // All 3 groups have 100% error rate, no recent diagnosis + }); + + test("should return false if there was a diagnosis tool call in the last 3 groups", () => { + const messages = [ + new AIMessage({ content: "AI message 1" }), // AI message 1 (part of last 3) + createToolMessage("1", "tool1", "error"), + createToolMessage("2", "tool2", "error"), + createToolMessage("3", "tool3", "error"), + createToolMessage("4", "tool4", "error"), // 100% error rate + + new AIMessage({ content: "AI message 2" }), // AI message 2 (part of last 3) + createToolMessage("5", "tool5", "error"), + createToolMessage("6", "tool6", "error"), + createToolMessage("7", "tool7", "error"), + createToolMessage("8", "diagnose_error", "success", true), // Diagnosis tool call + + new AIMessage({ content: "AI message 3" }), // AI message 3 (part of last 3) + createToolMessage("9", "tool9", "error"), + createToolMessage("10", "tool10", "error"), + createToolMessage("11", "tool11", "error"), // 100% error rate + ]; + + expect(shouldDiagnoseError(messages)).toBe(false); // Should not diagnose due to recent diagnosis + }); + + test("should return true if diagnosis tool call was more than 3 groups ago", () => { + const messages = [ + new AIMessage({ content: "AI message 1" }), // AI message 1 (NOT part of last 3) + createToolMessage("1", "tool1", "error"), + createToolMessage("2", "diagnose_error", "success", true), // Diagnosis tool call (old) + + new AIMessage({ content: "AI message 2" }), // AI message 2 (part of last 3) + createToolMessage("3", "tool3", "error"), + createToolMessage("4", "tool4", "error"), + createToolMessage("5", "tool5", "error"), + createToolMessage("6", "tool6", "success"), // 75% error rate + + new AIMessage({ content: "AI message 3" }), // AI message 3 (part of last 3) + createToolMessage("7", "tool7", "error"), + createToolMessage("8", "tool8", "error"), + createToolMessage("9", "tool9", "error"), // 100% error rate + + new AIMessage({ content: "AI message 4" }), // AI message 4 (part of last 3) + createToolMessage("10", "tool10", "error"), + createToolMessage("11", "tool11", "error"), + createToolMessage("12", "tool12", "success"), + createToolMessage("13", "tool13", "error"), // 75% error rate + ]; + + expect(shouldDiagnoseError(messages)).toBe(true); // Should diagnose since old diagnosis is outside last 3 groups + }); + + test("should return false if diagnosis tool call is in the most recent group", () => { + const messages = [ + new AIMessage({ content: "AI message 1" }), // AI message 1 (part of last 3) + createToolMessage("1", "tool1", "error"), + createToolMessage("2", "tool2", "error"), + createToolMessage("3", "tool3", "error"), // 100% error rate + + new AIMessage({ content: "AI message 2" }), // AI message 2 (part of last 3) + createToolMessage("4", "tool4", "error"), + createToolMessage("5", "tool5", "error"), + createToolMessage("6", "tool6", "error"), // 100% error rate + + new AIMessage({ content: "AI message 3" }), // AI message 3 (part of last 3) + createToolMessage("7", "tool7", "error"), + createToolMessage("8", "tool8", "error"), + createToolMessage("9", "tool9", "error"), + createToolMessage("10", "diagnose_error", "success", true), // Recent diagnosis + ]; + + expect(shouldDiagnoseError(messages)).toBe(false); // Should not diagnose due to recent diagnosis }); }); }); diff --git a/apps/open-swe/src/utils/tool-message-error.ts b/apps/open-swe/src/utils/tool-message-error.ts index 6e26690c..f96e4e33 100644 --- a/apps/open-swe/src/utils/tool-message-error.ts +++ b/apps/open-swe/src/utils/tool-message-error.ts @@ -63,27 +63,68 @@ export function calculateErrorRate(group: ToolMessage[]): number { return errorCount / group.length; } +/** + * Check if there was a diagnosis tool call within the last N tool message groups + * @param messages Array of messages to check + * @param groupCount Number of recent groups to check + * @returns True if a diagnosis tool call was found in the recent groups + */ +function hasRecentDiagnosisToolCall( + messages: Array, + groupCount: number, +): boolean { + const allGroups: ToolMessage[][] = []; + let currentGroup: ToolMessage[] = []; + let processingToolsForAI = false; + + for (let i = 0; i < messages.length; i++) { + const message = messages[i]; + + if (isAIMessage(message)) { + if (currentGroup.length > 0) { + allGroups.push([...currentGroup]); + currentGroup = []; + } + processingToolsForAI = true; + } else if (isToolMessage(message) && processingToolsForAI) { + currentGroup.push(message); + } else if (!isToolMessage(message) && processingToolsForAI) { + if (currentGroup.length > 0) { + allGroups.push([...currentGroup]); + currentGroup = []; + } + processingToolsForAI = false; + } + } + + if (currentGroup.length > 0) { + allGroups.push(currentGroup); + } + + const recentGroups = allGroups.slice(-groupCount); + return recentGroups.some((group) => + group.some((message) => message.additional_kwargs?.is_diagnosis), + ); +} + /** * Whether or not to route to the diagnose error step. This is true if: * - the last three tool call groups all have >= 75% error rates - * - * TBD: Should this be checking that each of the last 3 have >= 75% error rates, - * or >= 75% error rate of all tool messages from the last 3 groups? + * - there hasn't been a diagnose error tool call within the last three message groups * * @param messages All messages to analyze */ -export function shouldDiagnoseError(messages: Array) { - // Group tool messages by their parent AI message +export function shouldDiagnoseError(messages: Array) { const toolGroups = groupToolMessagesByAIMessage(messages); - // If we don't have at least 3 groups, we can't make a determination if (toolGroups.length < 3) return false; - // Get the last three groups const lastThreeGroups = toolGroups.slice(-3); - // Check if all of the last three groups have an error rate >= 75% - const ERROR_THRESHOLD = 0.75; // 75% + const hasRecentDiagnosis = hasRecentDiagnosisToolCall(messages, 3); + if (hasRecentDiagnosis) return false; + + const ERROR_THRESHOLD = 0.75; return lastThreeGroups.every( (group) => calculateErrorRate(group) >= ERROR_THRESHOLD, );