diff --git a/src/index.mjs b/src/index.mjs index ecb07e4..2c8118c 100644 --- a/src/index.mjs +++ b/src/index.mjs @@ -156,9 +156,22 @@ export function compact(payload, opts = {}) { for (let j = 0; j < Math.max(0, trimmableEnd); j++) { if (!isSystem(list[j])) { idx = j; break; } } if (idx < 0) break; // nothing left safe to drop const removed = list[idx]; - after -= blocksOf(removed.content, removed.role || "user") - .reduce((total, unit) => total + counter(unit.text), 0); - list.splice(idx, 1); + const pairedIndexes = new Set([idx]); + const removedBlocks = Array.isArray(removed.content) ? removed.content : []; + const removedToolUses = removedBlocks.filter((b) => b && b.type === "tool_use").map((b) => b.id).filter(Boolean); + if (removedToolUses.length) { + for (let k = 0; k < list.length; k++) { + if (k === idx) continue; + const blocks = Array.isArray(list[k].content) ? list[k].content : []; + if (blocks.some((b) => b && b.type === "tool_result" && removedToolUses.includes(b.tool_use_id))) pairedIndexes.add(k); + } + } + const removedMessages = [...pairedIndexes].sort((a, b) => b - a).map((index) => list[index]); + for (const message of removedMessages) { + after -= blocksOf(message.content, message.role || "user") + .reduce((total, unit) => total + counter(unit.text), 0); + } + for (const index of [...pairedIndexes].sort((a, b) => b - a)) list.splice(index, 1); actions.push("drop:oldest-message"); if (++i > 1000) break; } diff --git a/test/basic.test.mjs b/test/basic.test.mjs index 5c4158f..1a2c66c 100644 --- a/test/basic.test.mjs +++ b/test/basic.test.mjs @@ -172,3 +172,18 @@ test("units rejects non-array payload.messages with TypeError", () => { { name: "TypeError", message: /payload\.messages must be an array/ }, ); }); + +test("compact does not leave tool_result without its tool_use when trimming", () => { + const payload = { + messages: [ + { role: "assistant", content: [{ type: "tool_use", id: "call-1", input: {} }] }, + { role: "user", content: [{ type: "tool_result", tool_use_id: "call-1", content: "result" }] }, + ...Array.from({ length: 8 }, (_, i) => ({ role: "user", content: "old message " + i + " x".repeat(100) })), + ], + }; + const { payload: out } = compact(payload, { maxTokens: 20, keepLastTurns: 2 }); + const toolUses = new Set(out.messages.flatMap((m) => (Array.isArray(m.content) ? m.content : [])).filter((b) => b?.type === "tool_use").map((b) => b.id)); + for (const b of out.messages.flatMap((m) => (Array.isArray(m.content) ? m.content : []))) { + if (b?.type === "tool_result") assert.ok(toolUses.has(b.tool_use_id)); + } +});