fix(agent): handle excluded custom context messages

This commit is contained in:
Armin Ronacher
2026-06-14 17:29:15 +02:00
parent f4942bbc11
commit 6c75c67fb7
10 changed files with 338 additions and 22 deletions
+3 -2
View File
@@ -334,13 +334,14 @@ export class Agent {
await this.runPromptMessages(messages);
}
/** Continue from the current transcript. The last message must be a user or tool-result message. */
/** Continue from the current transcript. The last LLM-context message must be a user or tool-result message. */
async continue(): Promise<void> {
if (this.activeRun) {
throw new Error("Agent is already processing. Wait for completion before continuing.");
}
const lastMessage = this._state.messages[this._state.messages.length - 1];
const llmMessages = await this.convertToLlm(this._state.messages);
const lastMessage = llmMessages[llmMessages.length - 1];
if (!lastMessage) {
throw new Error("No messages to continue from");
}
@@ -269,10 +269,20 @@ export function estimateTokens(message: AgentMessage): number {
return 0;
}
function isExcludedCustomContextEntry(entry: SessionTreeEntry): boolean {
return (
(entry.type === "custom_message" && entry.excludeFromContext === true) ||
(entry.type === "message" && entry.message.role === "custom" && entry.message.excludeFromContext === true)
);
}
function findValidCutPoints(entries: SessionTreeEntry[], startIndex: number, endIndex: number): number[] {
const cutPoints: number[] = [];
for (let i = startIndex; i < endIndex; i++) {
const entry = entries[i];
if (isExcludedCustomContextEntry(entry)) {
continue;
}
switch (entry.type) {
case "message": {
const role = entry.message.role;
@@ -313,6 +323,9 @@ function findValidCutPoints(entries: SessionTreeEntry[], startIndex: number, end
export function findTurnStartIndex(entries: SessionTreeEntry[], entryIndex: number, startIndex: number): number {
for (let i = entryIndex; i >= startIndex; i--) {
const entry = entries[i];
if (isExcludedCustomContextEntry(entry)) {
continue;
}
if (entry.type === "branch_summary" || entry.type === "custom_message") {
return i;
}
@@ -371,6 +384,9 @@ export function findCutPoint(
if (prevEntry.type === "compaction") {
break;
}
if (isExcludedCustomContextEntry(prevEntry)) {
break;
}
if (prevEntry.type === "message") {
break;
}
+62 -2
View File
@@ -1,7 +1,19 @@
import { type AssistantMessage, type AssistantMessageEvent, EventStream, getModel } from "@earendil-works/pi-ai";
import {
type AssistantMessage,
type AssistantMessageEvent,
EventStream,
getModel,
type Message,
} from "@earendil-works/pi-ai";
import { Type } from "typebox";
import { describe, expect, it } from "vitest";
import { Agent, type AgentEvent, type AgentTool, type AgentToolUpdateCallback } from "../src/index.ts";
import {
Agent,
type AgentEvent,
type AgentMessage,
type AgentTool,
type AgentToolUpdateCallback,
} from "../src/index.ts";
// Mock stream that mimics AssistantMessageEventStream
class MockAssistantStream extends EventStream<AssistantMessageEvent, AssistantMessage> {
@@ -586,6 +598,54 @@ describe("Agent", () => {
expect(agent.state.messages[agent.state.messages.length - 1].role).toBe("assistant");
});
it("continue() should process queued follow-up messages after filtered state messages", async () => {
let providerCallCount = 0;
let providerMessages: Message[] = [];
const displayOnlyMessage = {
role: "displayOnly",
content: "status",
timestamp: Date.now(),
} as unknown as AgentMessage;
const agent = new Agent({
convertToLlm: (messages) =>
messages
.filter((message) => (message as { role: string }).role !== "displayOnly")
.filter(
(message) => message.role === "user" || message.role === "assistant" || message.role === "toolResult",
) as Message[],
streamFn: (_model, context) => {
providerCallCount++;
providerMessages = context.messages;
const stream = new MockAssistantStream();
queueMicrotask(() => {
stream.push({ type: "done", reason: "stop", message: createAssistantMessage("Processed") });
});
return stream;
},
});
agent.state.messages = [
{
role: "user",
content: [{ type: "text", text: "Initial" }],
timestamp: Date.now() - 10,
},
createAssistantMessage("Initial response"),
displayOnlyMessage,
];
agent.followUp({
role: "user",
content: [{ type: "text", text: "Queued follow-up" }],
timestamp: Date.now(),
});
await agent.continue();
expect(providerCallCount).toBe(1);
expect(providerMessages[providerMessages.length - 1]?.role).toBe("user");
expect(agent.state.messages).toContain(displayOnlyMessage);
});
it("continue() should keep one-at-a-time steering semantics from assistant tail", async () => {
let responseCount = 0;
const agent = new Agent({
@@ -234,6 +234,42 @@ describe("harness compaction", () => {
isSplitTurn: false,
});
const cutUser = createMessageEntry(createUserMessage("inspect file"));
const cutAssistant = createMessageEntry(
{
...createAssistantMessage("calling tool"),
content: [{ type: "toolCall", id: "call-2", name: "read", arguments: { path: "file.ts" } }],
},
cutUser.id,
);
const cutToolResult = createMessageEntry(
{
role: "toolResult",
toolCallId: "call-2",
toolName: "read",
content: [{ type: "text", text: "x".repeat(1000) }],
isError: false,
timestamp: Date.now(),
},
cutAssistant.id,
);
const excludedCustom: CustomMessageEntry = {
type: "custom_message",
id: createId(),
parentId: cutToolResult.id,
timestamp: new Date().toISOString(),
customType: "status",
content: "tool finished",
display: true,
excludeFromContext: true,
};
const cutAssistantFinal = createMessageEntry(createAssistantMessage("done"), excludedCustom.id);
expect(findCutPoint([cutUser, cutAssistant, cutToolResult, excludedCustom, cutAssistantFinal], 0, 5, 2)).toEqual({
firstKeptEntryIndex: 4,
turnStartIndex: 0,
isSplitTurn: true,
});
const user = createMessageEntry(createUserMessage("user"));
const compaction = createCompactionEntry("summary", user.id, user.id);
const assistant = createMessageEntry(createAssistantMessage("assistant"), compaction.id);
@@ -410,6 +446,56 @@ describe("harness compaction", () => {
expect(preparation?.messagesToSummarize).toEqual([]);
});
it("ignores excluded custom messages when finding split-turn prefixes", () => {
const user = createMessageEntry(createUserMessage("inspect file"));
const assistantWithToolCall = createMessageEntry(
{
...createAssistantMessage("calling tool"),
content: [{ type: "toolCall", id: "call-1", name: "read", arguments: { path: "file.ts" } }],
},
user.id,
);
const customMessage: CustomMessageEntry = {
type: "custom_message",
id: createId(),
parentId: assistantWithToolCall.id,
timestamp: new Date().toISOString(),
customType: "status",
content: "tool is running",
display: true,
excludeFromContext: true,
};
const toolResult = createMessageEntry(
{
role: "toolResult",
toolCallId: "call-1",
toolName: "read",
content: [{ type: "text", text: "x".repeat(1000) }],
isError: false,
timestamp: Date.now(),
},
customMessage.id,
);
const assistantFinal = createMessageEntry(createAssistantMessage("done"), toolResult.id);
const preparation = getOrThrow(
prepareCompaction([user, assistantWithToolCall, customMessage, toolResult, assistantFinal], {
enabled: true,
reserveTokens: 0,
keepRecentTokens: 1,
}),
);
expect(preparation).toBeDefined();
expect(preparation?.isSplitTurn).toBe(true);
expect(preparation?.firstKeptEntryId).toBe(assistantFinal.id);
expect(preparation?.turnPrefixMessages.map((message) => message.role)).toEqual([
"user",
"assistant",
"toolResult",
]);
});
it("skips excluded custom messages before branch summary token budgeting", () => {
const user = createMessageEntry(createUserMessage("keep"));
const customMessage: CustomMessageEntry = {
+29 -14
View File
@@ -583,6 +583,25 @@ export class AgentSession {
return undefined;
}
private _removeAssistantMessageFromState(target: AssistantMessage): void {
const messages = this.agent.state.messages;
for (let i = messages.length - 1; i >= 0; i--) {
const message = messages[i];
if (
message === target ||
(message.role === "assistant" &&
message.timestamp === target.timestamp &&
message.provider === target.provider &&
message.model === target.model &&
message.stopReason === target.stopReason &&
message.errorMessage === target.errorMessage)
) {
this.agent.state.messages = [...messages.slice(0, i), ...messages.slice(i + 1)];
return;
}
}
}
private _replaceMessageInPlace(target: AgentMessage, replacement: AgentMessage): void {
// Agent-core stores the finalized message object in its state before emitting message_end.
// SessionManager persistence happens later in _handleAgentEvent() with event.message.
@@ -1850,11 +1869,8 @@ export class AgentSession {
this._overflowRecoveryAttempted = true;
// Remove the error message from agent state (it IS saved to session for history,
// but we don't want it in context for the retry)
const messages = this.agent.state.messages;
if (messages.length > 0 && messages[messages.length - 1].role === "assistant") {
this.agent.state.messages = messages.slice(0, -1);
}
return await this._runAutoCompaction("overflow", true);
this._removeAssistantMessageFromState(assistantMessage);
return await this._runAutoCompaction("overflow", true, assistantMessage);
}
// Case 2: Threshold - context is getting large
@@ -1889,7 +1905,11 @@ export class AgentSession {
/**
* Internal: Run auto-compaction with events.
*/
private async _runAutoCompaction(reason: "overflow" | "threshold", willRetry: boolean): Promise<boolean> {
private async _runAutoCompaction(
reason: "overflow" | "threshold",
willRetry: boolean,
retryAssistantMessage?: AssistantMessage,
): Promise<boolean> {
const settings = this.settingsManager.getCompactionSettings();
this._emit({ type: "compaction_start", reason });
@@ -2037,10 +2057,8 @@ export class AgentSession {
this._emit({ type: "compaction_end", reason, result, aborted: false, willRetry });
if (willRetry) {
const messages = this.agent.state.messages;
const lastMsg = messages[messages.length - 1];
if (lastMsg?.role === "assistant" && (lastMsg as AssistantMessage).stopReason === "error") {
this.agent.state.messages = messages.slice(0, -1);
if (retryAssistantMessage) {
this._removeAssistantMessageFromState(retryAssistantMessage);
}
return true;
}
@@ -2525,10 +2543,7 @@ export class AgentSession {
});
// Remove error message from agent state (keep in session for history)
const messages = this.agent.state.messages;
if (messages.length > 0 && messages[messages.length - 1].role === "assistant") {
this.agent.state.messages = messages.slice(0, -1);
}
this._removeAssistantMessageFromState(message);
// Wait with exponential backoff (abortable)
this._retryAbortController = new AbortController();
@@ -313,10 +313,20 @@ export function estimateTokens(message: AgentMessage): number {
* and will be kept.
* BashExecutionMessage is treated like a user message (user-initiated context).
*/
function isExcludedCustomContextEntry(entry: SessionEntry): boolean {
return (
(entry.type === "custom_message" && entry.excludeFromContext === true) ||
(entry.type === "message" && entry.message.role === "custom" && entry.message.excludeFromContext === true)
);
}
function findValidCutPoints(entries: SessionEntry[], startIndex: number, endIndex: number): number[] {
const cutPoints: number[] = [];
for (let i = startIndex; i < endIndex; i++) {
const entry = entries[i];
if (isExcludedCustomContextEntry(entry)) {
continue;
}
switch (entry.type) {
case "message": {
const role = entry.message.role;
@@ -361,6 +371,9 @@ function findValidCutPoints(entries: SessionEntry[], startIndex: number, endInde
export function findTurnStartIndex(entries: SessionEntry[], entryIndex: number, startIndex: number): number {
for (let i = entryIndex; i >= startIndex; i--) {
const entry = entries[i];
if (isExcludedCustomContextEntry(entry)) {
continue;
}
// branch_summary and custom_message are user-role messages, can start a turn
if (entry.type === "branch_summary" || entry.type === "custom_message") {
return i;
@@ -444,6 +457,9 @@ export function findCutPoint(
if (prevEntry.type === "compaction") {
break;
}
if (isExcludedCustomContextEntry(prevEntry)) {
break;
}
if (prevEntry.type === "message") {
// Stop if we hit any message
break;
+12 -3
View File
@@ -128,10 +128,19 @@ export async function runPrintMode(runtimeHost: AgentSessionRuntime, options: Pr
if (mode === "text") {
const state = session.state;
const lastMessage = state.messages[state.messages.length - 1];
let assistantMsg: AssistantMessage | undefined;
for (let i = state.messages.length - 1; i >= 0; i--) {
const message = state.messages[i];
if (message.role === "custom" && message.excludeFromContext) {
continue;
}
if (message.role === "assistant") {
assistantMsg = message;
}
break;
}
if (lastMessage?.role === "assistant") {
const assistantMsg = lastMessage as AssistantMessage;
if (assistantMsg) {
if (assistantMsg.stopReason === "error" || assistantMsg.stopReason === "aborted") {
console.error(assistantMsg.errorMessage || `Request ${assistantMsg.stopReason}`);
exitCode = 1;
@@ -346,6 +346,31 @@ describe("findCutPoint", () => {
expect(result.turnStartIndex).toBe(2); // Turn 2 starts at index 2
}
});
it("should not select excluded custom messages as cut points", () => {
const user = createMessageEntry(createUserMessage("inspect file"));
const assistantWithToolCall = createMessageEntry({
...createAssistantMessage("calling tool"),
content: [{ type: "toolCall", id: "call-1", name: "read", arguments: { path: "file.ts" } }],
});
const toolResult = createMessageEntry({
role: "toolResult",
toolCallId: "call-1",
toolName: "read",
content: [{ type: "text", text: "x".repeat(1000) }],
isError: false,
timestamp: Date.now(),
});
const excludedCustom = createCustomMessageEntry("tool finished", true);
const assistantFinal = createMessageEntry(createAssistantMessage("done"));
const entries = [user, assistantWithToolCall, toolResult, excludedCustom, assistantFinal];
const result = findCutPoint(entries, 0, entries.length, 2);
expect(result.firstKeptEntryIndex).toBe(4);
expect(result.turnStartIndex).toBe(0);
expect(result.isSplitTurn).toBe(true);
});
});
describe("buildSessionContext", () => {
@@ -448,6 +473,39 @@ describe("prepareCompaction with custom messages", () => {
expect(preparation!.tokensBefore).toBe(1);
expect(preparation!.messagesToSummarize).toEqual([]);
});
it("should ignore excluded custom messages when finding split-turn prefixes", () => {
const user = createMessageEntry(createUserMessage("inspect file"));
const assistantWithToolCall = createMessageEntry({
...createAssistantMessage("calling tool"),
content: [{ type: "toolCall", id: "call-1", name: "read", arguments: { path: "file.ts" } }],
});
const excludedCustom = createCustomMessageEntry("tool is running", true);
const toolResult = createMessageEntry({
role: "toolResult",
toolCallId: "call-1",
toolName: "read",
content: [{ type: "text", text: "x".repeat(1000) }],
isError: false,
timestamp: Date.now(),
});
const assistantFinal = createMessageEntry(createAssistantMessage("done"));
const preparation = prepareCompaction([user, assistantWithToolCall, excludedCustom, toolResult, assistantFinal], {
enabled: true,
reserveTokens: 0,
keepRecentTokens: 1,
});
expect(preparation).toBeDefined();
expect(preparation!.isSplitTurn).toBe(true);
expect(preparation!.firstKeptEntryId).toBe(assistantFinal.id);
expect(preparation!.turnPrefixMessages.map((message) => message.role)).toEqual([
"user",
"assistant",
"toolResult",
]);
});
});
describe("prepareCompaction with previous compaction", () => {
@@ -514,12 +514,20 @@ describe("AgentSession queue characterization", () => {
it("delivers follow-ups queued during agent_end", async () => {
let sent = false;
let sawFollowUpInProvider = false;
const harness = await createHarness({
extensionFactories: [
(pi: ExtensionAPI) => {
pi.on("agent_end", async () => {
if (sent) return;
sent = true;
pi.sendMessage({
customType: "status",
content: "display only",
display: true,
details: {},
excludeFromContext: true,
});
pi.sendUserMessage("conflict report", { deliverAs: "followUp" });
});
},
@@ -527,11 +535,23 @@ describe("AgentSession queue characterization", () => {
});
harnesses.push(harness);
harness.setResponses([fauxAssistantMessage("reply"), fauxAssistantMessage("follow-up reply")]);
harness.setResponses([
fauxAssistantMessage("reply"),
(context) => {
sawFollowUpInProvider = context.messages.some(
(message) => message.role === "user" && getMessageText(message) === "conflict report",
);
return fauxAssistantMessage("follow-up reply");
},
]);
await harness.session.prompt("hello");
await harness.session.agent.waitForIdle();
expect(sawFollowUpInProvider).toBe(true);
expect(getUserTexts(harness)).toEqual(["hello", "conflict report"]);
expect(
harness.session.messages.some((message) => message.role === "custom" && message.customType === "status"),
).toBe(true);
});
});
@@ -52,6 +52,41 @@ describe("AgentSession retry and event characterization", () => {
expect(harness.session.isRetrying).toBe(false);
});
it("retries when an excluded custom message follows the transient error", async () => {
const harness = await createHarness({
settings: { retry: { enabled: true, maxRetries: 3, baseDelayMs: 1 } },
extensionFactories: [
(pi) => {
let sent = false;
pi.on("message_end", async (event) => {
if (sent || event.message.role !== "assistant" || event.message.stopReason !== "error") return;
sent = true;
pi.sendMessage({
customType: "status",
content: "display only",
display: true,
details: {},
excludeFromContext: true,
});
});
},
],
});
harnesses.push(harness);
harness.setResponses([
fauxAssistantMessage("", { stopReason: "error", errorMessage: "overloaded_error" }),
fauxAssistantMessage("recovered"),
]);
await harness.session.prompt("test");
expect(harness.faux.state.callCount).toBe(2);
expect(
harness.session.messages.some((message) => message.role === "custom" && message.customType === "status"),
).toBe(true);
expect(harness.session.messages[harness.session.messages.length - 1]?.role).toBe("assistant");
});
it("retries multiple transient failures and succeeds on the final attempt", async () => {
const harness = await createHarness({ settings: { retry: { enabled: true, maxRetries: 3, baseDelayMs: 1 } } });
harnesses.push(harness);