fix(agent): clear deferred choices on branches

A pre-model gate can defer a claimed hard tool choice for the next call. Branch transitions cleared the coding-agent queue but left that agent-owned value alive, allowing an obsolete forced tool to cross into the replacement transcript.

Expose the narrow deferred-choice reset at the Agent owner and invoke it from the shared session-scoped tool-state cleanup used by both branch paths. Failed session switches retain their existing rollback behavior.

Signed-off-by: Christian Stewart <christian@aperture.us>
This commit is contained in:
Christian Stewart
2026-07-26 01:26:08 -07:00
parent c5e5f6dbc3
commit 01ecab6df7
4 changed files with 120 additions and 47 deletions
+5 -1
View File
@@ -975,11 +975,15 @@ export class Agent {
this.#followUpQueue = [];
}
clearDeferredToolChoice() {
this.#deferredToolChoice = undefined;
}
clearAllQueues() {
this.#steeringQueue = [];
this.#followUpQueue = [];
this.#notifySteeringWaiters();
this.#deferredToolChoice = undefined;
this.clearDeferredToolChoice();
}
hasQueuedMessages(): boolean {
+70 -42
View File
@@ -98,9 +98,6 @@ describe("agentLoop with AgentMessage", () => {
expect(await stream.result()).toEqual([prompt]);
expect(mock.calls).toHaveLength(0);
expect(events.map(event => event.type)).toContain("agent_end");
expect(events.filter(event => event.type === "message_end")).toEqual([
expect.objectContaining({ message: prompt }),
]);
});
it("returns detailed telemetry when awaiting detailed() directly", async () => {
@@ -2518,12 +2515,6 @@ describe("agentLoop event-driven steering watch", () => {
const toolSchema = type({ value: "string" });
const tool: AgentTool<typeof toolSchema> = {
describe("agentLoop pre-model-call gate", () => {
const echoToolSchema = type({ value: "string" });
function echoTool(executed: string[]): AgentTool<typeof echoToolSchema> {
const toolSchema = echoToolSchema;
return {
name: "echo",
label: "Echo",
description: "Echo tool",
@@ -2880,6 +2871,76 @@ describe("agentLoop pre-model-call gate", () => {
};
const stream = agentLoop([createUserMessage("start")], context, config, undefined, mock.stream);
const drain = (async () => {
for await (const _event of stream) {
// drain
}
})();
const completed = await Promise.race([drain.then(() => true), Bun.sleep(1000).then(() => false)]);
try {
expect(completed).toBe(true);
expect(executed).toEqual(["only"]);
} finally {
check.resolve(false);
await drain;
}
});
it("stops watching after a steering subscription rejects", async () => {
let waitCalls = 0;
const toolSchema = type({ value: "string" });
const tool: AgentTool<typeof toolSchema> = {
name: "echo",
label: "Echo",
description: "Echo tool",
parameters: toolSchema,
concurrency: "exclusive",
async execute(_toolCallId, params) {
await Bun.sleep(0);
return { content: [{ type: "text", text: `ok:${params.value}` }], details: { value: params.value } };
},
};
const context: AgentContext = { systemPrompt: [""], messages: [], tools: [tool] };
const mock = createMockModel({
responses: [
{ content: [{ type: "toolCall", id: "tool-1", name: "echo", arguments: { value: "only" } }] },
{ content: ["done"] },
],
});
const config: AgentLoopConfig = {
model: mock.model,
convertToLlm: identityConverter,
interruptMode: "immediate",
hasSteeringMessages: () => ({ queued: false }),
waitForSteeringMessages: () => {
waitCalls++;
return Promise.reject(new Error("subscription unavailable"));
},
getSteeringMessages: async () => [],
};
const stream = agentLoop([createUserMessage("start")], context, config, undefined, mock.stream);
for await (const _event of stream) {
// drain
}
expect(waitCalls).toBe(1);
});
});
describe("agentLoop pre-model-call gate", () => {
const echoToolSchema = type({ value: "string" });
function echoTool(executed: string[]): AgentTool<typeof echoToolSchema> {
const toolSchema = echoToolSchema;
return {
name: "echo",
label: "Echo",
description: "Echo tool",
parameters: toolSchema,
concurrency: "exclusive",
async execute(_toolCallId, params) {
executed.push(params.value);
return { content: [{ type: "text", text: `ok:${params.value}` }], details: { value: params.value } };
},
};
@@ -3116,31 +3177,6 @@ describe("agentLoop pre-model-call gate", () => {
// drain
}
})();
const completed = await Promise.race([drain.then(() => true), Bun.sleep(1000).then(() => false)]);
try {
expect(completed).toBe(true);
expect(executed).toEqual(["only"]);
} finally {
check.resolve(false);
await drain;
}
});
it("stops watching after a steering subscription rejects", async () => {
let waitCalls = 0;
const toolSchema = type({ value: "string" });
const tool: AgentTool<typeof toolSchema> = {
name: "echo",
label: "Echo",
description: "Echo tool",
parameters: toolSchema,
concurrency: "exclusive",
async execute(_toolCallId, params) {
await Bun.sleep(0);
return { content: [{ type: "text", text: `ok:${params.value}` }], details: { value: params.value } };
},
};
const context: AgentContext = { systemPrompt: [""], messages: [], tools: [tool] };
await gateEntered.promise;
controller.abort();
await drain;
@@ -3416,13 +3452,6 @@ describe("agentLoop pre-model-call gate", () => {
const config: AgentLoopConfig = {
model: mock.model,
convertToLlm: identityConverter,
interruptMode: "immediate",
hasSteeringMessages: () => ({ queued: false }),
waitForSteeringMessages: () => {
waitCalls++;
return Promise.reject(new Error("subscription unavailable"));
},
getSteeringMessages: async () => [],
beforeModelCall: () => undefined,
};
@@ -3431,7 +3460,6 @@ describe("agentLoop pre-model-call gate", () => {
// drain
}
expect(waitCalls).toBe(1);
expect(executed).toEqual(["only"]);
});
});
@@ -4222,6 +4222,7 @@ export class AgentSession {
/** Drop mutable tool decisions and directives owned by the previous logical session. */
#clearSessionScopedToolState(): void {
this.agent.clearDeferredToolChoice();
this.#toolChoiceQueue.clear();
this.#tools.clearAcpPermissionDecisions();
}
@@ -1,7 +1,7 @@
import { afterEach, beforeEach, expect, it } from "bun:test";
import { afterEach, beforeEach, expect, it, vi } from "bun:test";
import * as path from "node:path";
import { Agent, type AgentTool } from "@oh-my-pi/pi-agent-core";
import { AssistantMessageEventStream } from "@oh-my-pi/pi-ai/utils/event-stream";
import { createMockModel, type MockModel } from "@oh-my-pi/pi-ai/providers/mock";
import { getBundledModel } from "@oh-my-pi/pi-catalog/models";
import { ModelRegistry } from "@oh-my-pi/pi-coding-agent/config/model-registry";
import { Settings } from "@oh-my-pi/pi-coding-agent/config/settings";
@@ -15,6 +15,8 @@ import { type } from "arktype";
let tempDir: TempDir;
let authStorage: AuthStorage | undefined;
let session: AgentSession;
let sessionManager: SessionManager;
let mock: MockModel;
beforeEach(async () => {
tempDir = TempDir.createSync("@pi-agent-session-force-tool-");
@@ -25,7 +27,7 @@ beforeEach(async () => {
authStorage.setRuntimeApiKey("anthropic", "test-key");
const modelRegistry = new ModelRegistry(authStorage, path.join(tempDir.path(), "models.yml"));
const settings = Settings.isolated({ "compaction.enabled": false });
const sessionManager = SessionManager.inMemory(tempDir.path());
sessionManager = SessionManager.inMemory(tempDir.path());
const emptyObjectSchema = type("object");
@@ -44,7 +46,10 @@ beforeEach(async () => {
execute: async () => ({ content: [{ type: "text" as const, text: "ok" }] }),
};
mock = createMockModel({ handler: () => ({ content: ["done"] }) });
const agent = new Agent({
getToolChoice: () => session.nextToolChoiceDirective(),
getApiKey: () => "test-key",
initialState: {
model,
@@ -53,7 +58,7 @@ beforeEach(async () => {
messages: [],
},
convertToLlm,
streamFn: () => new AssistantMessageEventStream(),
streamFn: mock.stream,
});
session = new AgentSession({
@@ -75,6 +80,14 @@ afterEach(async () => {
tempDir.removeSync();
});
async function deferForcedWrite(): Promise<void> {
session.setForcedToolChoice("write");
session.agent.setBeforeModelCall(() => ({ stop: true, reason: "session transition" }));
await session.agent.prompt("defer");
session.agent.setBeforeModelCall(undefined);
expect(mock.calls).toHaveLength(0);
}
it("forces specific tool, then transitions to none, then clears", () => {
session.setForcedToolChoice("write");
@@ -104,3 +117,30 @@ it("drops an unavailable forced choice with the rest of its sequence", async ()
it("throws when forcing a non-active tool", () => {
expect(() => session.setForcedToolChoice("read")).toThrow('Tool "read" is not currently active.');
});
it("drops a deferred forced choice when branching", async () => {
const entryId = sessionManager.appendMessage({
role: "user",
content: [{ type: "text", text: "branch target" }],
timestamp: Date.now(),
});
await deferForcedWrite();
await session.branch(entryId);
await session.agent.prompt("new branch");
expect(mock.calls).toHaveLength(1);
expect(mock.calls[0]?.options?.toolChoice).toBeUndefined();
});
it("retains a deferred forced choice when session switching rolls back", async () => {
await deferForcedWrite();
const failure = new Error("switch failed");
vi.spyOn(sessionManager, "setSessionFile").mockRejectedValueOnce(failure);
await expect(session.switchSession(path.join(tempDir.path(), "target.jsonl"))).rejects.toBe(failure);
await session.agent.prompt("retry current session");
expect(mock.calls).toHaveLength(1);
expect(mock.calls[0]?.options?.toolChoice).toEqual({ type: "tool", name: "write" });
});