Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 10 additions & 11 deletions src/CodexAcpClient.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,12 +2,16 @@ import {isCodexAuthRequest} from "./CodexAuthMethod";
import type {EmbeddedResourceResource} from "@agentclientprotocol/sdk";
import * as acp from "@agentclientprotocol/sdk";
import {type McpServer, RequestError} from "@agentclientprotocol/sdk";
import type {ApprovalHandler, CodexAppServerClient, ElicitationHandler} from "./CodexAppServerClient";
import type {
ApprovalHandler,
CodexAppServerClient,
ElicitationHandler,
McpStartupResult,
} from "./CodexAppServerClient";
import open from "open";
import type {Disposable} from "vscode-jsonrpc";
import type {
ClientInfo,
McpStartupCompleteEvent,
ReasoningEffort,
ServerNotification
} from "./app-server";
Expand Down Expand Up @@ -274,17 +278,12 @@ export class CodexAcpClient {
};
}

async awaitMcpStartup(mcpStartupVersion: number): Promise<Array<string>> {
const startup = await this.codexClient.awaitMcpStartup(mcpStartupVersion);
return startup.ready;
}

async awaitMcpStartupResult(mcpStartupVersion: number): Promise<McpStartupCompleteEvent> {
return await this.codexClient.awaitMcpStartup(mcpStartupVersion);
async awaitMcpServerStartup(serverNames: Array<string>, afterVersion: number): Promise<McpStartupResult> {
return await this.codexClient.awaitMcpServerStartup(serverNames, afterVersion);
}

getMcpStartupCompleteVersion(): number {
return this.codexClient.getMcpStartupCompleteVersion();
getMcpServerStartupVersion(): number {
return this.codexClient.getMcpServerStartupVersion();
}

private createSessionConfig(projectPath: string, mcpServers: Array<McpServer>): JsonObject {
Expand Down
73 changes: 42 additions & 31 deletions src/CodexAcpServer.ts
Original file line number Diff line number Diff line change
Expand Up @@ -10,8 +10,9 @@ import {CodexApprovalHandler} from "./CodexApprovalHandler";
import {CodexElicitationHandler} from "./CodexElicitationHandler";
import {CodexAuthMethods, type CodexAuthRequest} from "./CodexAuthMethod";
import {CodexAcpClient, type SessionMetadata, type SessionMetadataWithThread} from "./CodexAcpClient";
import type {McpStartupResult} from "./CodexAppServerClient";
import {ACPSessionConnection, type UpdateSessionEvent} from "./ACPSessionConnection";
import type {McpStartupCompleteEvent, InputModality, ReasoningEffort} from "./app-server";
import type {InputModality, ReasoningEffort} from "./app-server";
import type {
Account,
CollabAgentToolCallStatus,
Expand Down Expand Up @@ -55,6 +56,7 @@ export interface SessionState {

interface PendingMcpStartupSession {
requestedServers: Set<string>;
afterVersion: number;
}

export class CodexAcpServer implements acp.Agent {
Expand Down Expand Up @@ -143,7 +145,10 @@ export class CodexAcpServer implements acp.Agent {

async getOrCreateSession(request: acp.NewSessionRequest | acp.ResumeSessionRequest): Promise<[SessionId, SessionModelState, SessionModeState]> {
await this.checkAuthorization();
const mcpStartupVersion = this.codexAcpClient.getMcpStartupCompleteVersion();
const requestedMcpServers = request.mcpServers ?? [];
const mcpServerStartupVersion = requestedMcpServers.length > 0
? this.codexAcpClient.getMcpServerStartupVersion()
: null;

let sessionMetadata: SessionMetadata;
if ("sessionId" in request) {
Expand All @@ -156,7 +161,7 @@ export class CodexAcpServer implements acp.Agent {

const accountResponse = await this.runWithProcessCheck(() => this.codexAcpClient.getAccount());
const {sessionId, currentModelId, models} = sessionMetadata;
const sessionMcpServers = await this.resolveSessionMcpServers(request.mcpServers ?? [], mcpStartupVersion, "sessionId" in request);
const sessionMcpServers = this.resolveSessionMcpServers(requestedMcpServers, "sessionId" in request);
const currentModel = this.findCurrentModel(models, currentModelId);
const sessionState: SessionState = {
sessionId: sessionId,
Expand All @@ -175,12 +180,12 @@ export class CodexAcpServer implements acp.Agent {
}
this.sessions.set(sessionId, sessionState);

const requestedMcpServers = request.mcpServers ?? [];
if (requestedMcpServers.length > 0) {
if (requestedMcpServers.length > 0 && mcpServerStartupVersion !== null) {
this.pendingMcpStartupSessions.set(sessionId, {
requestedServers: new Set(requestedMcpServers.map(server => server.name)),
afterVersion: mcpServerStartupVersion,
});
this.publishMcpStartupStatusAsync(sessionId, mcpStartupVersion);
this.publishMcpStartupStatusAsync(sessionId);
}

this.publishAvailableCommandsAsync(sessionId);
Expand Down Expand Up @@ -355,7 +360,10 @@ export class CodexAcpServer implements acp.Agent {
thread: Thread;
}> {
await this.checkAuthorization();
const mcpStartupVersion = this.codexAcpClient.getMcpStartupCompleteVersion();
const requestedMcpServers = request.mcpServers ?? [];
const mcpServerStartupVersion = requestedMcpServers.length > 0
? this.codexAcpClient.getMcpServerStartupVersion()
: null;

logger.log(`Load existing session: ${request.sessionId}...`);
const sessionMetadata: SessionMetadataWithThread = await this.runWithProcessCheck(() =>
Expand All @@ -364,7 +372,7 @@ export class CodexAcpServer implements acp.Agent {

const accountResponse = await this.runWithProcessCheck(() => this.codexAcpClient.getAccount());
const {sessionId, currentModelId, models, thread} = sessionMetadata;
const sessionMcpServers = await this.resolveSessionMcpServers(request.mcpServers ?? [], mcpStartupVersion, true);
const sessionMcpServers = this.resolveSessionMcpServers(requestedMcpServers, true);
const currentModel = this.findCurrentModel(models, currentModelId);
const sessionState: SessionState = {
sessionId: sessionId,
Expand All @@ -383,12 +391,12 @@ export class CodexAcpServer implements acp.Agent {
};
this.sessions.set(sessionId, sessionState);

const requestedMcpServers = request.mcpServers ?? [];
if (requestedMcpServers.length > 0) {
if (requestedMcpServers.length > 0 && mcpServerStartupVersion !== null) {
this.pendingMcpStartupSessions.set(sessionId, {
requestedServers: new Set(requestedMcpServers.map(server => server.name)),
afterVersion: mcpServerStartupVersion,
});
this.publishMcpStartupStatusAsync(sessionId, mcpStartupVersion);
this.publishMcpStartupStatusAsync(sessionId);
}

await this.availableCommands.publish(sessionId);
Expand Down Expand Up @@ -633,11 +641,10 @@ export class CodexAcpServer implements acp.Agent {
return sessionState;
}

private async resolveSessionMcpServers(
private resolveSessionMcpServers(
mcpServers: Array<acp.McpServer>,
mcpStartupVersion: number,
recoverFromStartup: boolean,
): Promise<Array<string>> {
): Array<string> {
// Explicit MCP servers from the request are the primary source of truth for the session.
const requestedServerNames = getRequestedMcpServerNames(mcpServers);
if (requestedServerNames.length > 0) {
Expand All @@ -647,27 +654,31 @@ export class CodexAcpServer implements acp.Agent {
if (!recoverFromStartup) {
return [];
}
// loadSession/resumeSession may omit mcpServers; in that case recover the ready names
// from the startup event associated with this thread start/resume checkpoint.
logger.log("Recovering MCP servers from startup state...");
return await this.runWithProcessCheck(() => this.codexAcpClient.awaitMcpStartup(mcpStartupVersion));
// Without a thread-scoped startup completion event, loadSession/resumeSession can no longer
// recover omitted session MCP server names. Treat the session set as unknown unless ACP
// explicitly provided mcpServers in the request.
logger.log("Skipping MCP server recovery for load/resume without explicit mcpServers");
return [];
}

private publishMcpStartupStatusAsync(sessionId: string, mcpStartupVersion: number): void {
void this.doPublishMcpStartupStatus(sessionId, mcpStartupVersion);
private publishMcpStartupStatusAsync(sessionId: string): void {
void this.doPublishMcpStartupStatus(sessionId);
}

private async doPublishMcpStartupStatus(sessionId: string, mcpStartupVersion: number): Promise<void> {
private async doPublishMcpStartupStatus(sessionId: string): Promise<void> {
const pendingStartup = this.pendingMcpStartupSessions.get(sessionId);
if (!pendingStartup) {
return;
}

try {
const mcpStartup = await this.runWithProcessCheck(() => this.codexAcpClient.awaitMcpStartupResult(mcpStartupVersion));
const sessionState = this.sessions.get(sessionId);
const pendingStartup = this.pendingMcpStartupSessions.get(sessionId);
if (sessionState && pendingStartup) {
sessionState.sessionMcpServers = mcpStartup.ready.filter(serverName =>
pendingStartup.requestedServers.has(serverName)
);
}
await this.publishMcpStartupStatus(sessionId, mcpStartup, pendingStartup?.requestedServers);
const mcpStartup = await this.runWithProcessCheck(() =>
this.codexAcpClient.awaitMcpServerStartup(
Array.from(pendingStartup.requestedServers),
pendingStartup.afterVersion,
)
);
await this.publishMcpStartupStatus(sessionId, mcpStartup, pendingStartup.requestedServers);
} catch (err) {
logger.error(`Failed to publish MCP startup status for session ${sessionId}`, err);
} finally {
Expand All @@ -677,7 +688,7 @@ export class CodexAcpServer implements acp.Agent {

private async publishMcpStartupStatus(
sessionId: string,
mcpStartup: McpStartupCompleteEvent,
mcpStartup: McpStartupResult,
requestedServers?: Set<string>
): Promise<void> {
const filteredStartup = requestedServers
Expand Down
Loading
Loading