fix: Dont diagnose err after just diagnosing err (#527)

* fix: Dont diagnose err after just diagnosing err

* cr
This commit is contained in:
Brace Sproul 2025-07-24 16:23:02 -07:00 • committed by GitHub
parent 2fd4b94a2b
commit 5c7a7dcb49
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 137 additions and 21 deletions

View file

@ -157,23 +157,98 @@ describe("Error diagnosis logic", () => {
expect(shouldDiagnoseError(messages)).toBe(false); expect(shouldDiagnoseError(messages)).toBe(false);
}); });
test("should ignore diagnostic tool messages", () => { test("should ignore diagnostic tool messages when calculating error rates", () => {
const messages = [ const messages = [
new AIMessage({ content: "AI message 1" }), // AI message 1 new AIMessage({ content: "AI message 0" }), // AI message 0 (NOT part of last 3)
createToolMessage("1", "tool1", "error"), createToolMessage("0", "tool0", "error"),
createToolMessage("2", "tool2", "error", true), // Diagnostic (ignored) createToolMessage("1", "diagnose_error", "success", true), // Old diagnostic (ignored and outside last 3)
new AIMessage({ content: "AI message 2" }), // AI message 2 new AIMessage({ content: "AI message 1" }), // AI message 1 (part of last 3)
createToolMessage("3", "tool3", "error"), createToolMessage("2", "tool1", "error"),
createToolMessage("4", "tool4", "error"), createToolMessage("3", "tool2", "error"),
new AIMessage({ content: "AI message 3" }), // AI message 3 new AIMessage({ content: "AI message 2" }), // AI message 2 (part of last 3)
createToolMessage("5", "tool5", "error"), createToolMessage("4", "tool3", "error"),
createToolMessage("6", "tool6", "error"), createToolMessage("5", "tool4", "error"),
createToolMessage("7", "tool7", "error", true), // Diagnostic (ignored)
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
}); });
}); });
}); });

View file

@ -63,27 +63,68 @@ export function calculateErrorRate(group: ToolMessage[]): number {
return errorCount / group.length; 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: * 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 * - the last three tool call groups all have >= 75% error rates
* * - there hasn't been a diagnose error tool call within the last three message groups
* 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?
* *
* @param messages All messages to analyze * @param messages All messages to analyze
*/ */
export function shouldDiagnoseError(messages: Array<any>) { export function shouldDiagnoseError(messages: Array<BaseMessage>) {
// Group tool messages by their parent AI message
const toolGroups = groupToolMessagesByAIMessage(messages); 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; if (toolGroups.length < 3) return false;
// Get the last three groups
const lastThreeGroups = toolGroups.slice(-3); const lastThreeGroups = toolGroups.slice(-3);
// Check if all of the last three groups have an error rate >= 75% const hasRecentDiagnosis = hasRecentDiagnosisToolCall(messages, 3);
const ERROR_THRESHOLD = 0.75; // 75% if (hasRecentDiagnosis) return false;
const ERROR_THRESHOLD = 0.75;
return lastThreeGroups.every( return lastThreeGroups.every(
(group) => calculateErrorRate(group) >= ERROR_THRESHOLD, (group) => calculateErrorRate(group) >= ERROR_THRESHOLD,
); );