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:
@@ -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 {
|
||||
|
||||
@@ -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" });
|
||||
});
|
||||
|
||||
Reference in New Issue
Block a user