diff --git a/src/CodexAcpClient.ts b/src/CodexAcpClient.ts index 5fd4dbef..4e7f2aa4 100644 --- a/src/CodexAcpClient.ts +++ b/src/CodexAcpClient.ts @@ -356,6 +356,12 @@ export class CodexAcpClient { baseUrl: gatewayConfig.config.base_url, } : this.getNativeProviderConfig(); + logger.log("providers/list", { + providerId: OPENAI_PROVIDER_ID, + overrideActive: gatewayConfig !== null, + apiType: current.apiType, + baseUrl: current.baseUrl, + }); return [ { providerId: OPENAI_PROVIDER_ID, @@ -398,6 +404,12 @@ export class CodexAcpClient { baseUrl: request.baseUrl, headers: request.headers, }); + logger.log("providers/set applied", { + providerId: request.providerId, + apiType: request.apiType, + baseUrl: request.baseUrl, + headerNames: Object.keys(request.headers ?? {}), + }); } /** @@ -405,9 +417,24 @@ export class CodexAcpClient { * unknown provider id is idempotent success (RFD behavior ยง7). */ disableProvider(request: acp.DisableProviderRequest): void { + const overrideWasActive = this.gatewayConfig !== null; if (request.providerId === OPENAI_PROVIDER_ID) { this.gatewayConfig = null; } + const current = this.gatewayConfig + ? { + apiType: gatewayApiTypeFromConfig(this.gatewayConfig), + baseUrl: this.gatewayConfig.config.base_url, + } + : this.getNativeProviderConfig(); + logger.log("providers/disable applied", { + providerId: request.providerId, + knownProvider: request.providerId === OPENAI_PROVIDER_ID, + overrideWasActive, + overrideActive: this.gatewayConfig !== null, + restoredApiType: current.apiType, + restoredBaseUrl: current.baseUrl, + }); } async getAccount(): Promise { @@ -590,6 +617,19 @@ export class CodexAcpClient { mcpServers: Array ): Promise { const sessionRoots = [projectPath, ...additionalDirectories]; + const activeProvider = this.gatewayConfig + ? { + apiType: gatewayApiTypeFromConfig(this.gatewayConfig), + baseUrl: this.gatewayConfig.config.base_url, + } + : this.getNativeProviderConfig(); + logger.log("Creating session config", { + projectPath, + overrideActive: this.gatewayConfig !== null, + modelProvider: this.getModelProvider(), + apiType: activeProvider.apiType, + baseUrl: activeProvider.baseUrl, + }); const mergedConfig = { ...mergeGatewayConfig(this.config, this.gatewayConfig), projects: Object.fromEntries(sessionRoots.map(root => [root, { diff --git a/src/CodexAcpServer.ts b/src/CodexAcpServer.ts index 47e9373a..6e0c394b 100644 --- a/src/CodexAcpServer.ts +++ b/src/CodexAcpServer.ts @@ -7,11 +7,14 @@ import {type CodexAuthRequest, getCodexAuthMethods, isCodexAuthRequest} from "./ import {clientSupportsUrlElicitation} from "./ElicitationCapabilities"; import { CodexAcpClient, + type JsonObject, + OPENAI_PROVIDER_ID, type SessionMetadata, type SessionMetadataWithThread, type UrlElicitationRequester } from "./CodexAcpClient"; -import type {McpStartupResult} from "./CodexAppServerClient"; +import {CodexAppServerClient, type McpStartupResult} from "./CodexAppServerClient"; +import {type CodexConnection, startCodexConnection} from "./CodexJsonRpcConnection"; import {type AcpClientConnection, ACPSessionConnection, type UpdateSessionEvent} from "./ACPSessionConnection"; import type {InputModality, ReasoningEffort} from "./app-server"; import type {Account, Model, ReasoningEffortOption, Thread, ThreadGoal, ThreadItem, UserInput} from "./app-server/v2"; @@ -92,6 +95,7 @@ import { } from "./ContentChunks"; import {sameThreadGoalSnapshot, type ThreadGoalSnapshot, toThreadGoalSnapshot,} from "./ThreadGoalSnapshot"; import {randomUUID} from "node:crypto"; +import {once} from "node:events"; import { AIR_AGENT_FILE_CHANGE_REPORT_KEY, AIR_EXTENSION_CAPABILITIES_KEY, @@ -130,6 +134,7 @@ export interface SessionState { authProvider: string | null; cwd: string; additionalDirectories: string[]; + mcpServers?: Array; fastModeEnabled: boolean; currentModelSupportsFast: boolean; sessionMcpServers?: Array; @@ -213,6 +218,15 @@ interface ActivePrompt { complete: () => void; } +export interface CodexProcessState { + connection: CodexConnection; + codexPath: string | undefined; + config: JsonObject | undefined; + modelProvider: string | undefined; + stderr: string; + stderrProcess?: CodexConnection["process"]; +} + export class CodexAcpServer { private static readonly MODEL_NAME_TOKEN_OVERRIDES: Record = { gpt: "GPT", @@ -220,13 +234,13 @@ export class CodexAcpServer { codex: "Codex", }; - private readonly codexAcpClient: CodexAcpClient; + private codexAcpClient: CodexAcpClient; private readonly connection: AcpClientConnection; private readonly defaultAuthRequest: CodexAuthRequest | null; private readonly getExitCode: () => number | null; private readonly getRecentStderr: () => string; private readonly sessionFailureEpoch: string; - private readonly availableCommands: CodexCommands; + private availableCommands: CodexCommands; private clientInfo: acp.Implementation | null; private clientCapabilities: acp.ClientCapabilities | null; private terminalOutputMode: TerminalOutputMode; @@ -241,6 +255,9 @@ export class CodexAcpServer { private readonly sessionGenerations: Map; private readonly sessionOpenGenerations: Map; private readonly goalControlGenerations: Map; + private readonly codexProcessState: CodexProcessState | null; + private initializeRequest: acp.InitializeRequest | null = null; + private providerUpdate: Promise | null = null; constructor( connection: AcpClientConnection, @@ -248,6 +265,7 @@ export class CodexAcpServer { defaultAuthRequest?: CodexAuthRequest, getExitCode?: () => number | null, getRecentStderr?: () => string, + codexProcessState?: CodexProcessState, ) { this.sessions = new Map(); this.pendingMcpStartupSessions = new Map(); @@ -261,16 +279,22 @@ export class CodexAcpServer { this.connection = connection; this.codexAcpClient = codexAcpClient; this.defaultAuthRequest = defaultAuthRequest ?? null; - this.getExitCode = getExitCode ?? (() => null); - this.getRecentStderr = getRecentStderr ?? (() => ""); + this.codexProcessState = codexProcessState ?? null; + this.captureStderr(); + this.getExitCode = getExitCode ?? (() => this.codexProcessState?.connection.process.exitCode ?? null); + this.getRecentStderr = getRecentStderr ?? (() => this.codexProcessState?.stderr ?? ""); this.sessionFailureEpoch = randomUUID(); this.clientInfo = null; this.clientCapabilities = null; this.terminalOutputMode = "terminal_output_delta"; this.booleanConfigOptionsSupported = false; - this.availableCommands = new CodexCommands( - connection, - codexAcpClient, + this.availableCommands = this.createAvailableCommands(codexAcpClient); + } + + private createAvailableCommands(client: CodexAcpClient): CodexCommands { + return new CodexCommands( + this.connection, + client, (operation) => this.runWithProcessCheck(operation), () => this.refreshSessionsAuthState(null) ); @@ -282,6 +306,7 @@ export class CodexAcpServer { logger.log("Initialize request received"); this.clientInfo = _params.clientInfo ?? null; this.clientCapabilities = _params.clientCapabilities ?? null; + this.initializeRequest = _params; this.terminalOutputMode = resolveTerminalOutputMode(_params.clientCapabilities); this.booleanConfigOptionsSupported = clientSupportsBooleanConfigOptions(_params.clientCapabilities); await this.runWithProcessCheck(() => this.codexAcpClient.initialize(_params)); @@ -593,6 +618,7 @@ export class CodexAcpServer { authProvider: authProvider, cwd: request.cwd, additionalDirectories: sessionMetadata.additionalDirectories, + mcpServers: requestedMcpServers, fastModeEnabled: sessionMetadata.currentServiceTier === "fast", currentModelSupportsFast: currentModelSupportsFast, sessionMcpServers: sessionMcpServers, @@ -655,6 +681,9 @@ export class CodexAcpServer { } async loadSession(params: acp.LoadSessionRequest): Promise { + if (this.providerUpdate !== null) { + await this.providerUpdate; + } logger.log("Loading session...", {sessionId: params.sessionId}); const { sessionId, @@ -678,6 +707,9 @@ export class CodexAcpServer { } async resumeSession(params: acp.ResumeSessionRequest): Promise { + if (this.providerUpdate !== null) { + await this.providerUpdate; + } logger.log("Resuming session...", {sessionId: params.sessionId}); const [sessionId, modelState, modeState] = await this.getOrCreateSession(params); @@ -786,6 +818,9 @@ export class CodexAcpServer { async newSession( params: acp.NewSessionRequest, ): Promise { + if (this.providerUpdate !== null) { + await this.providerUpdate; + } logger.log("Starting new session..."); const [sessionId, modelState, modeState] = await this.getOrCreateSession(params); @@ -843,16 +878,115 @@ export class CodexAcpServer { return { providers: this.codexAcpClient.listProviders() }; } - setProvider(params: acp.SetProviderRequest): acp.SetProviderResponse { + async setProvider(params: acp.SetProviderRequest): Promise { this.codexAcpClient.setProvider(params); + await this.enqueueProviderUpdate((client) => client.setProvider(params)); return { }; } - disableProvider(params: acp.DisableProviderRequest): acp.DisableProviderResponse { + async disableProvider(params: acp.DisableProviderRequest): Promise { this.codexAcpClient.disableProvider(params); + if (params.providerId !== OPENAI_PROVIDER_ID) { + return { }; + } + await this.enqueueProviderUpdate((client) => client.disableProvider(params)); return { }; } + private async enqueueProviderUpdate(apply: (client: CodexAcpClient) => void): Promise { + const previous = this.providerUpdate?.catch(() => undefined) ?? Promise.resolve(); + const update = previous.then(async () => { + if (this.sessions.size === 0) { + return; + } + + const activePrompts = [...this.activePrompts.values()].map(prompt => prompt.completion); + if (activePrompts.length > 0) { + logger.log("Waiting for active prompts before provider restart", {count: activePrompts.length}); + await Promise.all(activePrompts); + } + + logger.log("Restarting Codex app-server for provider update", {sessionCount: this.sessions.size}); + const replacement = await this.restartCodexClient(); + apply(replacement); + if (this.initializeRequest === null) { + throw new Error("Cannot restart Codex app-server before ACP initialization"); + } + await replacement.initialize(this.initializeRequest); + this.codexAcpClient = replacement; + this.availableCommands = this.createAvailableCommands(replacement); + + const resumeErrors: unknown[] = []; + for (const session of this.sessions.values()) { + try { + await replacement.resumeSession({ + sessionId: session.sessionId, + cwd: session.cwd, + additionalDirectories: session.additionalDirectories, + mcpServers: session.mcpServers ?? [], + }); + session.authProvider = replacement.getModelProvider(); + logger.log("Resumed session after provider restart", {sessionId: session.sessionId}); + } catch (error) { + resumeErrors.push(error); + logger.error(`Failed to resume session ${session.sessionId} after provider restart`, error); + } + } + if (resumeErrors.length > 0) { + throw new AggregateError(resumeErrors, `Failed to resume ${resumeErrors.length} session(s) after provider restart`); + } + }); + this.providerUpdate = update; + try { + await update; + } finally { + if (this.providerUpdate === update) { + this.providerUpdate = null; + } + } + } + + private captureStderr(): void { + const state = this.codexProcessState; + if (state === null || state.stderrProcess === state.connection.process) { + return; + } + state.stderrProcess = state.connection.process; + state.connection.process.stderr.addListener("data", (data: Buffer) => { + state.stderr = (state.stderr + data.toString()).slice(-2 * 1024); + }); + } + + private async restartCodexClient(): Promise { + const state = this.codexProcessState; + if (state === null) { + throw new Error("Codex process state is unavailable"); + } + + const previous = state.connection; + const exited = previous.process.exitCode === null + ? once(previous.process, "exit") + : Promise.resolve(); + previous.process.stdin.end(); + const forceKill = setTimeout(() => { + if (previous.process.exitCode === null) { + logger.log("Codex still running 2s after provider restart; terminating process"); + previous.process.kill(); + } + }, 2000); + await exited; + clearTimeout(forceKill); + + state.stderr = ""; + state.connection = startCodexConnection(state.codexPath); + this.captureStderr(); + return new CodexAcpClient( + new CodexAppServerClient(state.connection.connection), + state.config, + state.modelProvider, + ); + } + private async refreshSessionsAuthState(authProvider: string | null): Promise { if (this.sessions.size === 0) return; @@ -1482,6 +1616,7 @@ export class CodexAcpServer { authProvider: authProvider, cwd: request.cwd, additionalDirectories: sessionMetadata.additionalDirectories, + mcpServers: requestedMcpServers, fastModeEnabled: sessionMetadata.currentServiceTier === "fast", currentModelSupportsFast: currentModelSupportsFast, sessionMcpServers: sessionMcpServers, @@ -2092,6 +2227,9 @@ export class CodexAcpServer { signal?: AbortSignal, onTurnStarted?: () => void, ): Promise { + if (this.providerUpdate !== null) { + await this.providerUpdate; + } logger.log("Prompt received", { sessionId: params.sessionId, prompt: params.prompt, diff --git a/src/Logger.ts b/src/Logger.ts index b64fdec6..fd588777 100644 --- a/src/Logger.ts +++ b/src/Logger.ts @@ -18,6 +18,7 @@ class Logger { try { fs.mkdirSync(logDir, {recursive: true}); this.logFilePath = path.join(logDir, "app-server.log"); + this.log("Logger initialized", {logFilePath: this.logFilePath}); } catch (ex) { console.error("Failed to initialize logger directory", ex); this.logFilePath = null; @@ -32,7 +33,7 @@ class Logger { if (!this.logFilePath) return; try { const timestamp = this.formatTimestamp(new Date()); - const serializedContext = context ? ` ${JSON.stringify(context)}` : ""; + const serializedContext = ` ${JSON.stringify({pid: process.pid, ...context})}`; if (!message.startsWith('[')) message = `[SYS] ${message}`; const line = `${timestamp} ${message}${serializedContext}`; diff --git a/src/__tests__/CodexACPAgent/providers.test.ts b/src/__tests__/CodexACPAgent/providers.test.ts index 113c63fd..45bc744c 100644 --- a/src/__tests__/CodexACPAgent/providers.test.ts +++ b/src/__tests__/CodexACPAgent/providers.test.ts @@ -1,15 +1,10 @@ import {describe, expect, it, vi} from "vitest"; import * as acp from "@agentclientprotocol/sdk"; -import {createCodexMockTestFixture} from "../acp-test-utils"; +import {createCodexMockTestFixture, createTestSessionState} from "../acp-test-utils"; import {CodexAcpClient, CUSTOM_GATEWAY_PROVIDER_ID, OPENAI_PROVIDER_ID} from "../../CodexAcpClient"; -function expectInvalidParams(fn: () => unknown): void { - let caught: unknown; - try { - fn(); - } catch (err) { - caught = err; - } +async function expectInvalidParams(fn: () => unknown): Promise { + const caught = await Promise.resolve().then(fn).catch((err: unknown) => err); expect(caught).toBeInstanceOf(acp.RequestError); expect((caught as acp.RequestError).code).toBe(-32602); } @@ -82,30 +77,30 @@ describe("Configurable LLM providers (providers/*)", () => { expect(JSON.stringify(provider)).not.toContain("super-secret"); }); - it("rejects an unsupported apiType with invalid_params", () => { + it("rejects an unsupported apiType with invalid_params", async () => { const fixture = createCodexMockTestFixture(); const agent = fixture.getCodexAcpAgent(); - expectInvalidParams(() => agent.setProvider({ + await expectInvalidParams(() => agent.setProvider({ providerId: OPENAI_PROVIDER_ID, apiType: "anthropic", baseUrl: "https://example.com", })); }); - it("rejects an unknown providerId with invalid_params", () => { + it("rejects an unknown providerId with invalid_params", async () => { const fixture = createCodexMockTestFixture(); const agent = fixture.getCodexAcpAgent(); - expectInvalidParams(() => agent.setProvider({ + await expectInvalidParams(() => agent.setProvider({ providerId: "does-not-exist", apiType: "openai", baseUrl: "https://example.com", })); }); - it("rejects a malformed baseUrl with invalid_params", () => { + it("rejects a malformed baseUrl with invalid_params", async () => { const fixture = createCodexMockTestFixture(); const agent = fixture.getCodexAcpAgent(); - expectInvalidParams(() => agent.setProvider({ + await expectInvalidParams(() => agent.setProvider({ providerId: OPENAI_PROVIDER_ID, apiType: "openai", baseUrl: " ", @@ -173,6 +168,132 @@ describe("Configurable LLM providers (providers/*)", () => { })); }); + it("keeps native session creation unchanged when the providers API is unused", async () => { + const restart = vi.fn(); + const fixture = createCodexMockTestFixture(restart); + const agent = fixture.getCodexAcpAgent(); + const codexAcpClient = fixture.getCodexAcpClient(); + const codexAppServerClient = fixture.getCodexAppServerClient(); + + vi.spyOn(codexAcpClient, "authRequired").mockResolvedValue(false); + const threadStartSpy = vi.spyOn(codexAppServerClient, "threadStart") + .mockRejectedValue(new Error("stop after capturing config")); + + await expect(agent.newSession({cwd: "/workspace", mcpServers: []})).rejects.toThrow(); + + expect(restart).not.toHaveBeenCalled(); + expect(threadStartSpy).toHaveBeenCalledOnce(); + expect(threadStartSpy).toHaveBeenCalledWith(expect.objectContaining({ + modelProvider: null, + config: expect.not.objectContaining({ + model_providers: expect.objectContaining({ + [CUSTOM_GATEWAY_PROVIDER_ID]: expect.anything(), + }), + }), + })); + }); + + it("restarts app-server and resumes all loaded sessions through the full provider lifecycle", async () => { + const firstGatewayReplacement = createCodexMockTestFixture().getCodexAcpClient(); + const secondGatewayReplacement = createCodexMockTestFixture().getCodexAcpClient(); + const nativeReplacement = createCodexMockTestFixture().getCodexAcpClient(); + vi.spyOn(firstGatewayReplacement, "initialize").mockResolvedValue(); + const firstGatewayResume = vi.spyOn(firstGatewayReplacement, "resumeSession").mockResolvedValue({} as never); + vi.spyOn(secondGatewayReplacement, "initialize").mockResolvedValue(); + const secondGatewayResume = vi.spyOn(secondGatewayReplacement, "resumeSession").mockResolvedValue({} as never); + vi.spyOn(nativeReplacement, "initialize").mockResolvedValue(); + const nativeResume = vi.spyOn(nativeReplacement, "resumeSession").mockResolvedValue({} as never); + const restart = vi.fn() + .mockResolvedValueOnce(firstGatewayReplacement) + .mockResolvedValueOnce(secondGatewayReplacement) + .mockResolvedValueOnce(nativeReplacement); + const fixture = createCodexMockTestFixture(restart); + const agent = fixture.getCodexAcpAgent(); + await agent.initialize({protocolVersion: acp.PROTOCOL_VERSION}); + const sessions = (agent as unknown as {sessions: Map>}).sessions; + sessions.set("thread-1", createTestSessionState({sessionId: "thread-1", cwd: "/one"})); + sessions.set("thread-2", createTestSessionState({sessionId: "thread-2", cwd: "/two"})); + + await agent.setProvider({ + providerId: OPENAI_PROVIDER_ID, + apiType: "openai", + baseUrl: "https://gateway.example/v1", + headers: {Authorization: "Bearer secret"}, + }); + + expect(restart).toHaveBeenCalledTimes(1); + expect(firstGatewayReplacement.getModelProvider()).toBe(CUSTOM_GATEWAY_PROVIDER_ID); + expect(firstGatewayResume).toHaveBeenCalledTimes(2); + expect(firstGatewayResume).toHaveBeenCalledWith(expect.objectContaining({sessionId: "thread-1", cwd: "/one"})); + expect(firstGatewayResume).toHaveBeenCalledWith(expect.objectContaining({sessionId: "thread-2", cwd: "/two"})); + + await agent.setProvider({ + providerId: OPENAI_PROVIDER_ID, + apiType: "openai", + baseUrl: "https://second-gateway.example/v1", + headers: {Authorization: "Bearer replacement"}, + }); + + expect(restart).toHaveBeenCalledTimes(2); + expect(secondGatewayReplacement.getModelProvider()).toBe(CUSTOM_GATEWAY_PROVIDER_ID); + expect(secondGatewayResume).toHaveBeenCalledTimes(2); + expect(agent.listProviders({}).providers[0]!.current).toEqual({ + apiType: "openai", + baseUrl: "https://second-gateway.example/v1", + }); + + await agent.disableProvider({providerId: OPENAI_PROVIDER_ID}); + + expect(restart).toHaveBeenCalledTimes(3); + expect(nativeReplacement.getModelProvider()).toBeNull(); + expect(nativeResume).toHaveBeenCalledTimes(2); + expect(agent.listProviders({}).providers[0]!.current).toEqual({ + apiType: "openai", + baseUrl: "https://api.openai.com/v1", + }); + }); + + it("attempts every session resume and allows a later provider update after one fails", async () => { + const failedReplacement = createCodexMockTestFixture().getCodexAcpClient(); + const recoveredReplacement = createCodexMockTestFixture().getCodexAcpClient(); + vi.spyOn(failedReplacement, "initialize").mockResolvedValue(); + const failedResume = vi.spyOn(failedReplacement, "resumeSession") + .mockRejectedValueOnce(new Error("thread-1 failed")) + .mockResolvedValue({} as never); + vi.spyOn(recoveredReplacement, "initialize").mockResolvedValue(); + const recoveredResume = vi.spyOn(recoveredReplacement, "resumeSession").mockResolvedValue({} as never); + const restart = vi.fn() + .mockResolvedValueOnce(failedReplacement) + .mockResolvedValueOnce(recoveredReplacement); + const fixture = createCodexMockTestFixture(restart); + const agent = fixture.getCodexAcpAgent(); + await agent.initialize({protocolVersion: acp.PROTOCOL_VERSION}); + const sessions = (agent as unknown as {sessions: Map>}).sessions; + sessions.set("thread-1", createTestSessionState({sessionId: "thread-1", cwd: "/one"})); + sessions.set("thread-2", createTestSessionState({sessionId: "thread-2", cwd: "/two"})); + + await expect(agent.setProvider({ + providerId: OPENAI_PROVIDER_ID, + apiType: "openai", + baseUrl: "https://broken-gateway.example/v1", + })).rejects.toThrow("Failed to resume 1 session(s)"); + + expect(failedResume).toHaveBeenCalledTimes(2); + + await expect(agent.setProvider({ + providerId: OPENAI_PROVIDER_ID, + apiType: "openai", + baseUrl: "https://recovered-gateway.example/v1", + })).resolves.toEqual({}); + + expect(restart).toHaveBeenCalledTimes(2); + expect(recoveredResume).toHaveBeenCalledTimes(2); + expect(agent.listProviders({}).providers[0]!.current).toEqual({ + apiType: "openai", + baseUrl: "https://recovered-gateway.example/v1", + }); + }); + it("shares state with the legacy gateway auth method", async () => { const fixture = createCodexMockTestFixture(); const codexAcpClient = fixture.getCodexAcpClient(); diff --git a/src/__tests__/acp-test-utils.ts b/src/__tests__/acp-test-utils.ts index de992ccd..8d6c9071 100644 --- a/src/__tests__/acp-test-utils.ts +++ b/src/__tests__/acp-test-utils.ts @@ -98,7 +98,13 @@ export function createBaseTestFixture(config: ConnectionConfig): TestFixture { const codexAppServerClient = new CodexAppServerClient(config.connection); const codexAcpClient = new CodexAcpClient(codexAppServerClient); - const codexAcpAgent = new CodexAcpServer(acpConnection, codexAcpClient, undefined, config.getExitCode); + const codexAcpAgent = new CodexAcpServer( + acpConnection, + codexAcpClient, + undefined, + config.getExitCode, + undefined, + ); const transportEvents: CodexConnectionEvent[] = []; const codexEventHandlers: ((event: CodexConnectionEvent) => void)[] = []; @@ -255,7 +261,9 @@ export interface CodexMockTestFixture extends TestFixture { * Provides `sendServerRequest()` to simulate server-initiated requests (e.g., approval requests). * Provides `setPermissionResponse()` to control ACP permission dialog responses. */ -export function createCodexMockTestFixture(): CodexMockTestFixture { +export function createCodexMockTestFixture( + restartCodexClient?: () => Promise, +): CodexMockTestFixture { let unhandledNotificationHandler: ((notification: any) => void) | null = null; const requestHandlers = new Map Promise>(); @@ -309,6 +317,10 @@ export function createCodexMockTestFixture(): CodexMockTestFixture { eventHandlers: acpEventHandlers, } }); + if (restartCodexClient) { + vi.spyOn(baseFixture.getCodexAcpAgent() as any, "restartCodexClient") + .mockImplementation(restartCodexClient); + } return { ...baseFixture, diff --git a/src/index.ts b/src/index.ts index 0bddc4e2..15ad2256 100644 --- a/src/index.ts +++ b/src/index.ts @@ -3,7 +3,7 @@ import * as acp from "@agentclientprotocol/sdk"; import {z} from "zod"; import {startCodexConnection} from "./CodexJsonRpcConnection"; -import {CodexAcpServer} from "./CodexAcpServer"; +import {CodexAcpServer, type CodexProcessState} from "./CodexAcpServer"; import {createJsonStream} from "./StdUtils"; import {isCodexAuthRequest} from "./CodexAuthMethod"; import {CodexAcpClient} from "./CodexAcpClient"; @@ -88,21 +88,21 @@ function startAcpServer() { defaultAuthRequest: defaultAuthRequest ?? null, }); - const codexConnection = startCodexConnection(codexPath); - - const maxStderrTailChars = 2 * 1024; - let stderr = ""; - codexConnection.process.stderr.addListener("data", (data: Buffer) => { - stderr = (stderr + data.toString()).slice(-maxStderrTailChars); - }); + const codexProcessState: CodexProcessState = { + connection: startCodexConnection(codexPath), + codexPath, + config, + modelProvider, + stderr: "", + }; process.stdin.on("close", () => { - codexConnection.process.stdin.end(); + codexProcessState.connection.process.stdin.end(); // Kill the codex process if it doesn't exit naturally setTimeout(() => { - if (!codexConnection.process.killed) { + if (!codexProcessState.connection.process.killed) { logger.log("Codex still running 2s after stdin closed; terminating process"); - codexConnection.process.kill(); + codexProcessState.connection.process.kill(); } }, 2000); }); @@ -110,9 +110,16 @@ function startAcpServer() { const acpJsonStream = createJsonStream(process.stdin, process.stdout); function createAgent(connection: acp.AgentContext): CodexAcpServer { - const appServerClient = new CodexAppServerClient(codexConnection.connection); + const appServerClient = new CodexAppServerClient(codexProcessState.connection.connection); const codexClient = new CodexAcpClient(appServerClient, config, modelProvider); - return new CodexAcpServer(connection, codexClient, defaultAuthRequest, () => codexConnection.process.exitCode, () => stderr); + return new CodexAcpServer( + connection, + codexClient, + defaultAuthRequest, + undefined, + undefined, + codexProcessState, + ); } let codexAcpServer: CodexAcpServer | null = null;