diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index 0ac6a6c0e..415a1be94 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -2,6 +2,10 @@ ## [Unreleased] +### Fixed + +- Pending resolve reminders now trigger as soon as a preview action is queued, before the next assistant turn, with regression coverage in `agent-session-resolve-reminder` tests + ## [13.4.0] - 2026-03-01 ### Breaking Changes diff --git a/packages/coding-agent/src/session/agent-session.ts b/packages/coding-agent/src/session/agent-session.ts index 8450c2751..afc80bde4 100644 --- a/packages/coding-agent/src/session/agent-session.ts +++ b/packages/coding-agent/src/session/agent-session.ts @@ -291,6 +291,7 @@ export class AgentSession { // Event subscription state #unsubscribeAgent?: () => void; + #unsubscribePendingActionPush?: () => void; #eventListeners: AgentSessionEventListener[] = []; /** Tracks pending steering messages for UI display. Removed when delivered. */ @@ -397,6 +398,21 @@ export class AgentSession { this.#obfuscator = config.obfuscator; this.agent.providerSessionState = this.#providerSessionState; this.#pendingActionStore = config.pendingActionStore; + this.#unsubscribePendingActionPush = this.#pendingActionStore?.subscribePush(action => { + const reminderText = [ + "", + "This is a preview. Call the `resolve` tool to apply or discard these changes.", + "", + ].join("\n"); + this.agent.steer({ + role: "custom", + customType: "resolve-reminder", + content: reminderText, + display: false, + details: { toolName: action.sourceToolName }, + timestamp: Date.now(), + }); + }); this.#syncTodoPhasesFromBranch(); // Always subscribe to agent events for internal handling @@ -688,22 +704,6 @@ export class AgentSession { { deliverAs: "nextTurn" }, ); } - if (!isError && this.#pendingActionStore?.hasPending) { - const reminderText = [ - "", - "This is a preview. Call the `resolve` tool to apply or discard these changes.", - "", - ].join("\n"); - await this.sendCustomMessage( - { - customType: "resolve-reminder", - content: reminderText, - display: false, - details: { toolName }, - }, - { deliverAs: "nextTurn" }, - ); - } } } @@ -1443,6 +1443,8 @@ export class AgentSession { state.close(); } this.#providerSessionState.clear(); + this.#unsubscribePendingActionPush?.(); + this.#unsubscribePendingActionPush = undefined; this.#disconnectFromAgent(); this.#eventListeners = []; } diff --git a/packages/coding-agent/src/tools/pending-action.ts b/packages/coding-agent/src/tools/pending-action.ts index 315edba7f..71446eca0 100644 --- a/packages/coding-agent/src/tools/pending-action.ts +++ b/packages/coding-agent/src/tools/pending-action.ts @@ -10,9 +10,14 @@ export interface PendingAction { export class PendingActionStore { #actions: PendingAction[] = []; + #pushListeners = new Set<(action: PendingAction, count: number) => void>(); push(action: PendingAction): void { this.#actions.push(action); + const count = this.#actions.length; + for (const listener of this.#pushListeners) { + listener(action, count); + } } peek(): PendingAction | null { @@ -23,10 +28,21 @@ export class PendingActionStore { return this.#actions.pop() ?? null; } + subscribePush(listener: (action: PendingAction, count: number) => void): () => void { + this.#pushListeners.add(listener); + return () => { + this.#pushListeners.delete(listener); + }; + } + clear(): void { this.#actions = []; } + get count(): number { + return this.#actions.length; + } + get hasPending(): boolean { return this.#actions.length > 0; } diff --git a/packages/coding-agent/test/agent-session-resolve-reminder.test.ts b/packages/coding-agent/test/agent-session-resolve-reminder.test.ts new file mode 100644 index 000000000..8f635c044 --- /dev/null +++ b/packages/coding-agent/test/agent-session-resolve-reminder.test.ts @@ -0,0 +1,118 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; +import * as fs from "node:fs"; +import * as os from "node:os"; +import * as path from "node:path"; +import { Agent } from "@oh-my-pi/pi-agent-core"; +import { type AssistantMessage, getBundledModel } from "@oh-my-pi/pi-ai"; +import { AssistantMessageEventStream } from "@oh-my-pi/pi-ai/utils/event-stream"; +import { ModelRegistry } from "@oh-my-pi/pi-coding-agent/config/model-registry"; +import { Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; +import { AgentSession } from "@oh-my-pi/pi-coding-agent/session/agent-session"; +import { AuthStorage } from "@oh-my-pi/pi-coding-agent/session/auth-storage"; +import { SessionManager } from "@oh-my-pi/pi-coding-agent/session/session-manager"; +import { PendingActionStore } from "@oh-my-pi/pi-coding-agent/tools/pending-action"; +import { Snowflake } from "@oh-my-pi/pi-utils"; + +class MockAssistantStream extends AssistantMessageEventStream {} + +function createAssistantMessage(text: string): AssistantMessage { + return { + role: "assistant", + content: [{ type: "text", text }], + api: "anthropic-messages", + provider: "anthropic", + model: "mock", + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "stop", + timestamp: Date.now(), + }; +} + +describe("AgentSession resolve reminder", () => { + let session: AgentSession; + let tempDir: string; + let pendingActionStore: PendingActionStore; + let streamCallCount = 0; + + beforeEach(async () => { + tempDir = path.join(os.tmpdir(), `pi-resolve-reminder-test-${Snowflake.next()}`); + fs.mkdirSync(tempDir, { recursive: true }); + pendingActionStore = new PendingActionStore(); + streamCallCount = 0; + + const model = getBundledModel("anthropic", "claude-sonnet-4-5"); + if (!model) { + throw new Error("Test model not found in registry"); + } + + const authStorage = await AuthStorage.create(path.join(tempDir, "testauth.db")); + authStorage.setRuntimeApiKey("anthropic", "test-key"); + const modelRegistry = new ModelRegistry(authStorage, path.join(tempDir, "models.yml")); + + const agent = new Agent({ + initialState: { + model, + systemPrompt: "Test", + tools: [], + messages: [], + }, + streamFn: () => { + streamCallCount += 1; + if (streamCallCount === 1) { + pendingActionStore.push({ + label: "AST Edit: 1 replacement in 1 file", + sourceToolName: "ast_edit", + apply: async () => ({ content: [{ type: "text", text: "Applied" }] }), + }); + } + const stream = new MockAssistantStream(); + queueMicrotask(() => { + stream.push({ type: "start", partial: createAssistantMessage("") }); + stream.push({ type: "done", reason: "stop", message: createAssistantMessage("Done") }); + }); + return stream; + }, + }); + + session = new AgentSession({ + agent, + sessionManager: SessionManager.inMemory(), + settings: Settings.isolated(), + modelRegistry, + pendingActionStore, + }); + }); + + afterEach(async () => { + await session.dispose(); + if (fs.existsSync(tempDir)) { + fs.rmSync(tempDir, { recursive: true, force: true }); + } + vi.restoreAllMocks(); + }); + + it("forces an immediate steering turn and injects resolve reminder before second assistant response", async () => { + await session.prompt("run preview"); + + expect(streamCallCount).toBe(2); + + const messages = session.agent.state.messages; + const assistantIndices = messages + .map((message, index) => (message.role === "assistant" ? index : -1)) + .filter(index => index >= 0); + const reminderIndex = messages.findIndex( + message => message.role === "custom" && message.customType === "resolve-reminder", + ); + + expect(assistantIndices.length).toBe(2); + expect(reminderIndex).toBeGreaterThan(assistantIndices[0]!); + expect(reminderIndex).toBeLessThan(assistantIndices[1]!); + }); +});