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);
});
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
});
});
});

View file

@ -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,
);