diff --git a/src/CodexAcpClient.ts b/src/CodexAcpClient.ts index 459046ff..39939ce2 100644 --- a/src/CodexAcpClient.ts +++ b/src/CodexAcpClient.ts @@ -9,6 +9,7 @@ import type { SetDefaultModelParams, SetDefaultModelResponse } from "./app-server"; +import type {TurnCompletedNotification } from "./app-server/v2"; import type {JsonValue} from "./app-server/serde_json/JsonValue"; import type {Model} from "./app-server/v2"; import {ModelId} from "./ModelId"; @@ -101,7 +102,10 @@ export class CodexAcpClient { }; } - async sendPrompt(request: acp.PromptRequest, eventHandler: (result: ServerNotification) => void): Promise { + async sendPrompt( + request: acp.PromptRequest, + eventHandler: (result: ServerNotification) => void + ): Promise { this.codexClient.onServerNotification(request.sessionId, eventHandler); const input = request.prompt.filter(b => b.type === "text") @@ -119,13 +123,22 @@ export class CodexAcpClient { model: null, }); - await this.codexClient.awaitTurnCompleted(); + // Wait for turn completion + // If turnInterrupt() was called, Codex will send turn/completed event with status "interrupted" + return await this.codexClient.awaitTurnCompleted(); } async setModel(params: SetDefaultModelParams): Promise { return this.codexClient.setModelRequest(params); } + async turnInterrupt(params: { threadId: string, turnId: string }): Promise { + await this.codexClient.turnInterrupt({ + threadId: params.threadId, + turnId: params.turnId + }); + } + private async fetchAvailableModels(): Promise { const models: Model[] = []; let cursor: string | null = null; diff --git a/src/CodexAcpServer.ts b/src/CodexAcpServer.ts index 6d0679fc..128e6fc7 100644 --- a/src/CodexAcpServer.ts +++ b/src/CodexAcpServer.ts @@ -10,7 +10,7 @@ import {ModelId} from "./ModelId"; export interface SessionState { sessionMetadata: SessionMetadata; - pendingPrompt: AbortController | null; + currentTurnId: string | null; } export class CodexAcpServer implements acp.Agent { @@ -62,7 +62,7 @@ export class CodexAcpServer implements acp.Agent { const {sessionId, currentModelId, models} = sessionMetadata; this.sessions.set(sessionId, { sessionMetadata: sessionMetadata, - pendingPrompt: null + currentTurnId: null }); const availableModels = this.buildAvailableModels(models); @@ -153,25 +153,38 @@ export class CodexAcpServer implements acp.Agent { async prompt(params: acp.PromptRequest): Promise { const sessionState = this.getSessionState(params.sessionId); - sessionState.pendingPrompt?.abort(); - sessionState.pendingPrompt = new AbortController(); + sessionState.currentTurnId = null; try { const messageHandler = new CodexEventHandler(this.connection, sessionState); - await this.runWithProcessCheck(() => this.codexAcpClient.sendPrompt(params, (event) => messageHandler.handleNotification(event))); - } catch (err) { - if (sessionState.pendingPrompt.signal.aborted) { - return {stopReason: "cancelled"}; + const turnCompleted = await this.runWithProcessCheck(() => this.codexAcpClient.sendPrompt(params, (event) => messageHandler.handleNotification(event))); + + // Check if turn was interrupted (cancelled) + if (turnCompleted.turn.status === "interrupted") { + await this.connection.sessionUpdate({ + sessionId: params.sessionId, + update: { + sessionUpdate: "agent_message_chunk", + content: { + type: "text", + text: "*Conversation interrupted*" + } + } + }); + return { + stopReason: "cancelled", + }; } + return { + stopReason: "end_turn", + }; + } catch (err) { + console.error(`Prompt for session ${params.sessionId} failed:`, err); throw err; + } finally { + sessionState.currentTurnId = null; } - - sessionState.pendingPrompt = null; - - return { - stopReason: "end_turn", - }; } private async runWithProcessCheck(operation: () => Promise): Promise { @@ -191,7 +204,27 @@ export class CodexAcpServer implements acp.Agent { } async cancel(params: acp.CancelNotification): Promise { - //TODO not supported yet - this.sessions.get(params.sessionId)?.pendingPrompt?.abort(); + const sessionState = this.sessions.get(params.sessionId); + if (!sessionState) { + console.info(`Can not cancel: session ${params.sessionId} not found`); + return; + } + + if (!sessionState.currentTurnId) { + console.info(`Can not cancel: session ${params.sessionId} has no current turn`); + return; + } + + console.info(`Cancel session ${params.sessionId}, currentTurnId: ${sessionState.currentTurnId}...`); + try { + // After turnInterrupt(), Codex will send turn/completed event, which will naturally complete awaitTurnCompleted() + await this.codexAcpClient.turnInterrupt({ + threadId: params.sessionId, + turnId: sessionState.currentTurnId + }); + console.log(`Cancel - turnInterrupt succeeded`); + } catch (err) { + console.error(`Cancel - turnInterrupt failed:`, err); + } } } \ No newline at end of file diff --git a/src/CodexAppServerClient.ts b/src/CodexAppServerClient.ts index f4f763f2..fa8b44d6 100644 --- a/src/CodexAppServerClient.ts +++ b/src/CodexAppServerClient.ts @@ -13,6 +13,8 @@ import type { ThreadStartParams, ThreadStartResponse, TurnCompletedNotification, + TurnInterruptParams, + TurnInterruptResponse, TurnStartParams, TurnStartResponse } from "./app-server/v2"; @@ -45,6 +47,10 @@ export class CodexAppServerClient { return await this.sendRequest({ method: "turn/start", params: params }); } + async turnInterrupt(params: TurnInterruptParams): Promise { + return await this.sendRequest({ method: "turn/interrupt", params: params }); + } + async threadStart(params: ThreadStartParams): Promise { return await this.sendRequest({ method: "thread/start", params: params }); } diff --git a/src/CodexEventHandler.ts b/src/CodexEventHandler.ts index 19119005..88d7a473 100644 --- a/src/CodexEventHandler.ts +++ b/src/CodexEventHandler.ts @@ -53,12 +53,16 @@ export class CodexEventHandler { return await this.updatePlan(notification.params); case "error": return await this.createErrorEvent(notification.params); + case "turn/started": + this.sessionState.currentTurnId = notification.params.turn.id; + return null; + case "turn/completed": + this.sessionState.currentTurnId = null; + return null; case "item/reasoning/summaryTextDelta": //TODO streaming reasoning? case "item/reasoning/summaryPartAdded": //skipped events case "item/reasoning/textDelta": //for raw output - case "turn/started": - case "turn/completed": case "turn/diff/updated": case "item/commandExecution/outputDelta": case "item/fileChange/outputDelta": diff --git a/src/__tests__/CodexACPAgent/CodexAcpClient.test.ts b/src/__tests__/CodexACPAgent/CodexAcpClient.test.ts index 57cf8e2c..0ce998a3 100644 --- a/src/__tests__/CodexACPAgent/CodexAcpClient.test.ts +++ b/src/__tests__/CodexACPAgent/CodexAcpClient.test.ts @@ -77,10 +77,15 @@ describe('ACP server test', { timeout: 40_000 }, () => { const codexAcpAgent = fixture.getCodexAcpAgent(); - fixture.getCodexAppServerClient().turnStart = vi.fn().mockResolvedValue(undefined); - fixture.getCodexAppServerClient().awaitTurnCompleted = vi.fn().mockResolvedValue(undefined); + fixture.getCodexAppServerClient().turnStart = vi.fn().mockResolvedValue({ + turn: { id: "turn-id", items: [], status: "inProgress", error: null } + }); + fixture.getCodexAppServerClient().awaitTurnCompleted = vi.fn().mockResolvedValue({ + threadId: "id", + turn: { id: "turn-id", items: [], status: "completed", error: null } + }); const sessionState: SessionState = { - pendingPrompt: null, + currentTurnId: null, sessionMetadata: { sessionId: "id", currentModelId: "model-id", @@ -99,11 +104,16 @@ describe('ACP server test', { timeout: 40_000 }, () => { const mockFixture = createCodexMockTestFixture(); const codexAcpAgent = mockFixture.getCodexAcpAgent(); - mockFixture.getCodexAppServerClient().turnStart = vi.fn().mockResolvedValue(undefined); - mockFixture.getCodexAppServerClient().awaitTurnCompleted = vi.fn().mockResolvedValue(undefined); + mockFixture.getCodexAppServerClient().turnStart = vi.fn().mockResolvedValue({ + turn: { id: "turn-id", items: [], status: "inProgress", error: null } + }); + mockFixture.getCodexAppServerClient().awaitTurnCompleted = vi.fn().mockResolvedValue({ + threadId: "id", + turn: { id: "turn-id", items: [], status: "completed", error: null } + }); const sessionState: SessionState = { - pendingPrompt: null, + currentTurnId: null, sessionMetadata: { sessionId: "id", currentModelId: "model-id", @@ -143,11 +153,16 @@ describe('ACP server test', { timeout: 40_000 }, () => { const mockFixture = createCodexMockTestFixture(); const codexAcpAgent = mockFixture.getCodexAcpAgent(); - mockFixture.getCodexAppServerClient().turnStart = vi.fn().mockResolvedValue(undefined); - mockFixture.getCodexAppServerClient().awaitTurnCompleted = vi.fn().mockResolvedValue(undefined); + mockFixture.getCodexAppServerClient().turnStart = vi.fn().mockResolvedValue({ + turn: { id: "turn-id", items: [], status: "inProgress", error: null } + }); + mockFixture.getCodexAppServerClient().awaitTurnCompleted = vi.fn().mockResolvedValue({ + threadId: "id", + turn: { id: "turn-id", items: [], status: "completed", error: null } + }); const sessionState1: SessionState = { - pendingPrompt: null, + currentTurnId: null, sessionMetadata: { sessionId: "session-1", currentModelId: "model-id", @@ -155,7 +170,7 @@ describe('ACP server test', { timeout: 40_000 }, () => { } }; const sessionState2: SessionState = { - pendingPrompt: null, + currentTurnId: null, sessionMetadata: { sessionId: "session-2", currentModelId: "model-id",