diff --git a/packages/agent/src/agent.ts b/packages/agent/src/agent.ts index 278918f37..b9beeaa50 100644 --- a/packages/agent/src/agent.ts +++ b/packages/agent/src/agent.ts @@ -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 { diff --git a/packages/agent/test/agent-loop.test.ts b/packages/agent/test/agent-loop.test.ts index d347b27a7..264879224 100644 --- a/packages/agent/test/agent-loop.test.ts +++ b/packages/agent/test/agent-loop.test.ts @@ -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 = { -describe("agentLoop pre-model-call gate", () => { - const echoToolSchema = type({ value: "string" }); - - function echoTool(executed: string[]): AgentTool { - 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 = { + 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 { + 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 = { - 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"]); }); }); diff --git a/packages/coding-agent/src/session/agent-session.ts b/packages/coding-agent/src/session/agent-session.ts index 969bc457e..916185283 100644 --- a/packages/coding-agent/src/session/agent-session.ts +++ b/packages/coding-agent/src/session/agent-session.ts @@ -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(); } diff --git a/packages/coding-agent/test/agent-session-force-tool-choice.test.ts b/packages/coding-agent/test/agent-session-force-tool-choice.test.ts index 6aedd33b1..1f55264c8 100644 --- a/packages/coding-agent/test/agent-session-force-tool-choice.test.ts +++ b/packages/coding-agent/test/agent-session-force-tool-choice.test.ts @@ -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 { + 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" }); +});