mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 15:03:16 +00:00
fix: Dont diagnose err after just diagnosing err (#527)
* fix: Dont diagnose err after just diagnosing err * cr
This commit is contained in:
parent
2fd4b94a2b
commit
5c7a7dcb49
2 changed files with 137 additions and 21 deletions
|
|
@ -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
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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<BaseMessage>,
|
||||
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<any>) {
|
||||
// Group tool messages by their parent AI message
|
||||
export function shouldDiagnoseError(messages: Array<BaseMessage>) {
|
||||
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,
|
||||
);
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue