diff --git a/src/CodexAcpClient.ts b/src/CodexAcpClient.ts index 893dc85d..68b74e00 100644 --- a/src/CodexAcpClient.ts +++ b/src/CodexAcpClient.ts @@ -64,6 +64,9 @@ import { createUnavailableAgentFileChangeReport, } from "./AgentFileChangeReport"; import {CodexSubagentSubscriptions} from "./subagents/CodexSubagentSubscriptions"; +import {forkSession as runForkSession} from "./SessionFork"; +import type {SessionMetadata, SessionMetadataWithThread} from "./SessionMetadata"; +export type {SessionMetadata, SessionMetadataWithThread} from "./SessionMetadata"; /** * Well-known provider id for the client-configurable custom LLM gateway. @@ -487,6 +490,21 @@ export class CodexAcpClient { } } + async forkSession(request: acp.ForkSessionRequest): Promise { + const additionalDirectories = readAdditionalDirectories(request.cwd, request.additionalDirectories, request._meta); + return await runForkSession(request, additionalDirectories, { + codexClient: this.codexClient, + refreshSkills: (cwd, directories) => this.refreshSkills(cwd, directories), + createSessionConfig: (cwd, directories, mcpServers) => + this.createSessionConfig(cwd, directories, mcpServers), + getResumeModelProvider: () => this.getResumeModelProvider(), + fetchAvailableModels: () => this.fetchAvailableModels(), + createCurrentModelId: (models, model, reasoningEffort) => + this.createModelId(models, model, reasoningEffort).toString(), + getCollaborationMode: sessionId => this.getCollaborationMode(sessionId), + }); + } + async loadSession(request: acp.LoadSessionRequest, onSubscribed?: () => void): Promise { const additionalDirectories = readAdditionalDirectories(request.cwd, request.additionalDirectories, request._meta); await this.refreshSkills(request.cwd, additionalDirectories); @@ -1223,20 +1241,6 @@ class AgentFileChangeReportBudget { export type JsonObject = { [key in string]?: JsonValue } -export type SessionMetadata = { - sessionId: string, - currentModelId: string, - models: Model[], - collaborationMode: ModeKind, - modelProvider?: string | null, - currentServiceTier?: ServiceTier | null, - additionalDirectories: string[], -} - -export type SessionMetadataWithThread = SessionMetadata & { - thread: Thread, -} - function buildPromptItems(prompt: acp.ContentBlock[]): UserInput[] { return prompt.map((block): UserInput | null => { switch (block.type) { diff --git a/src/CodexAcpServer.ts b/src/CodexAcpServer.ts index f94978e4..2c7937f3 100644 --- a/src/CodexAcpServer.ts +++ b/src/CodexAcpServer.ts @@ -315,6 +315,7 @@ export class CodexAcpServer { list: { }, close: { }, delete: { }, + fork: { }, additionalDirectories: {}, subagents: {}, }; @@ -553,9 +554,13 @@ export class CodexAcpServer { return generation; } - async tryCreateSession(request: acp.NewSessionRequest | acp.ResumeSessionRequest): Promise<[SessionId, LegacySessionModelState, SessionModeState]> { - const requestedSessionGeneration = "sessionId" in request - ? this.beginSessionOpen(request.sessionId) + async tryCreateSession( + request: acp.NewSessionRequest | acp.ResumeSessionRequest | acp.ForkSessionRequest, + operation: "new" | "resume" | "fork" = "sessionId" in request ? "resume" : "new", + ): Promise<[SessionId, LegacySessionModelState, SessionModeState]> { + const existingSessionRequest = request as acp.ResumeSessionRequest | acp.ForkSessionRequest; + const requestedSessionGeneration = operation === "resume" + ? this.beginSessionOpen(existingSessionRequest.sessionId) : null; await this.checkAuthorization(); const requestedMcpServers = request.mcpServers ?? []; @@ -565,23 +570,28 @@ export class CodexAcpServer { let sessionMetadata: SessionMetadata; let resumeSubscribed = false; - if ("sessionId" in request) { - logger.log(`Resume existing session: ${request.sessionId}...`); + if (operation === "resume") { + const resumeRequest = request as acp.ResumeSessionRequest; + logger.log(`Resume existing session: ${resumeRequest.sessionId}...`); try { sessionMetadata = await this.runWithProcessCheck(() => - this.codexAcpClient.resumeSession(request, () => { + this.codexAcpClient.resumeSession(resumeRequest, () => { resumeSubscribed = true; }) ); } catch (err) { if (resumeSubscribed && requestedSessionGeneration !== null) { - await this.cleanupStaleSessionOpen(request.sessionId, requestedSessionGeneration); + await this.cleanupStaleSessionOpen(resumeRequest.sessionId, requestedSessionGeneration); } throw err; } + } else if (operation === "fork") { + const forkRequest = request as acp.ForkSessionRequest; + logger.log(`Fork existing session: ${forkRequest.sessionId}...`); + sessionMetadata = await this.runWithProcessCheck(() => this.codexAcpClient.forkSession(forkRequest)); } else { logger.log(`Create new session...`); - sessionMetadata = await this.runWithProcessCheck(() => this.codexAcpClient.newSession(request)); + sessionMetadata = await this.runWithProcessCheck(() => this.codexAcpClient.newSession(request as acp.NewSessionRequest)); } const {sessionId, currentModelId, models} = sessionMetadata; @@ -600,7 +610,7 @@ export class CodexAcpServer { resumeSubscribed = false; await this.closeStaleSessionOpen(sessionId, sessionGeneration); } - const sessionMcpServers = this.resolveSessionMcpServers(requestedMcpServers, "sessionId" in request); + const sessionMcpServers = this.resolveSessionMcpServers(requestedMcpServers, operation === "resume"); const currentModel = this.findCurrentModel(models, currentModelId); const currentModelSupportsFast = modelSupportsFast(currentModel); const sessionState: SessionState = { @@ -628,7 +638,7 @@ export class CodexAcpServer { terminalOutputMode: this.terminalOutputMode, goalRevision: 0, sessionTitle: null, - sessionTitleSource: "sessionId" in request ? "unknown" : "unset", + sessionTitleSource: operation === "resume" ? "unknown" : "unset", subagents: new CodexSubagentEventRouter( sessionId, clientSupportsSubagents(this.clientCapabilities), @@ -638,7 +648,8 @@ export class CodexAcpServer { this.sessions.set(sessionId, sessionState); resumeSubscribed = false; - if (requestedMcpServers.length > 0 && mcpServerStartupVersion !== null) { + const canPublishSessionUpdates = operation !== "fork"; + if (canPublishSessionUpdates && requestedMcpServers.length > 0 && mcpServerStartupVersion !== null) { this.pendingMcpStartupSessions.set(sessionId, { requestedServers: new Set(getRequestedMcpServerNames(requestedMcpServers)), afterVersion: mcpServerStartupVersion, @@ -646,8 +657,10 @@ export class CodexAcpServer { this.publishMcpStartupStatusAsync(sessionId); } - this.publishAvailableCommandsAsync(sessionState, sessionGeneration); - if ("sessionId" in request) { + if (canPublishSessionUpdates) { + this.publishAvailableCommandsAsync(sessionState, sessionGeneration); + } + if (operation === "resume") { this.publishCurrentGoalAsync(sessionState, sessionGeneration); } const sessionModelState: LegacySessionModelState = this.createModelState(models, currentModelId); @@ -733,6 +746,26 @@ export class CodexAcpServer { }; } + async forkSession(params: acp.ForkSessionRequest): Promise { + if (this.providerUpdate !== null) { + await this.providerUpdate; + } + logger.log("Forking session...", {sessionId: params.sessionId}); + try { + const [sessionId, , modeState] = await this.tryCreateSession(params, "fork"); + logger.log("Session forked", {sourceSessionId: params.sessionId, sessionId}); + return { + sessionId, + modes: modeState, + ...this.createSessionConfigOptionsResponse(this.getSessionState(sessionId)), + }; + } catch (e) { + const error = e instanceof Error ? e : new Error(String(e)); + await this.handleError(error); + throw e; + } + } + async listSessions(params: acp.ListSessionsRequest): Promise { logger.log("Listing sessions...", {cwd: params.cwd, cursor: params.cursor}); await this.checkAuthorization(); diff --git a/src/SessionFork.ts b/src/SessionFork.ts new file mode 100644 index 00000000..b52b7cae --- /dev/null +++ b/src/SessionFork.ts @@ -0,0 +1,134 @@ +import {createHash} from "node:crypto"; +import type * as acp from "@agentclientprotocol/sdk"; +import {RequestError} from "@agentclientprotocol/sdk"; +import type {CodexAppServerClient} from "./CodexAppServerClient"; +import type {ModeKind} from "./app-server/ModeKind"; +import type {ServiceTier} from "./app-server/ServiceTier"; +import type {Model, ThreadForkParams} from "./app-server/v2"; +import type {SessionMetadata} from "./SessionMetadata"; + +export type SessionForkDependencies = { + codexClient: CodexAppServerClient; + refreshSkills(cwd: string, additionalDirectories: string[]): Promise; + createSessionConfig( + cwd: string, + additionalDirectories: string[], + mcpServers: acp.McpServer[], + ): Promise>; + getResumeModelProvider(): Promise; + fetchAvailableModels(): Promise; + createCurrentModelId(models: Model[], model: string, reasoningEffort: string | null): string; + getCollaborationMode(sessionId: string): ModeKind; +}; + +export async function forkSession( + request: acp.ForkSessionRequest, + additionalDirectories: string[], + dependencies: SessionForkDependencies, +): Promise { + await dependencies.refreshSkills(request.cwd, additionalDirectories); + const lastTurnId = await resolveForkTurnId(request, dependencies.codexClient); + const response = await dependencies.codexClient.threadFork({ + config: await dependencies.createSessionConfig( + request.cwd, + additionalDirectories, + request.mcpServers ?? [], + ), + cwd: request.cwd, + ...(lastTurnId !== undefined && {lastTurnId}), + modelProvider: await dependencies.getResumeModelProvider(), + threadId: request.sessionId, + }); + await dependencies.codexClient.threadUnsubscribe({threadId: response.thread.id}); + + const models = await dependencies.fetchAvailableModels(); + return { + sessionId: response.thread.id, + currentModelId: dependencies.createCurrentModelId(models, response.model, response.reasoningEffort), + models, + collaborationMode: dependencies.getCollaborationMode(response.thread.id), + modelProvider: response.modelProvider, + currentServiceTier: response.serviceTier as ServiceTier ?? null, + additionalDirectories, + }; +} + +async function resolveForkTurnId( + request: acp.ForkSessionRequest, + codexClient: CodexAppServerClient, +): Promise { + const forkPoint = readAirForkPoint(request._meta); + if (!forkPoint) return undefined; + + const history = await codexClient.threadRead({ + threadId: request.sessionId, + includeTurns: true, + }); + const candidateIds = airForkMessageIdCandidates(forkPoint.messageId); + const itemTurnId = candidateIds + .map(candidateId => history.thread.turns.find(turn => turn.items.some(item => item.id === candidateId))?.id) + .find(turnId => turnId !== undefined); + if (itemTurnId) return itemTurnId; + + if (forkPoint.messageFingerprint) { + const matchingTurns = history.thread.turns.flatMap(turn => turn.items + .filter(item => item.type === "agentMessage" + && fingerprintAgentMessage(item.text) === forkPoint.messageFingerprint) + .map(() => turn.id)); + const fingerprintTurnId = matchingTurns[forkPoint.messageOccurrence - 1]; + if (fingerprintTurnId) return fingerprintTurnId; + } + + throw RequestError.invalidParams( + {messageId: forkPoint.messageId}, + `Fork point message ${forkPoint.messageId} was not found in session ${request.sessionId}`, + ); +} + +type AirForkPoint = { + messageId: string; + messageFingerprint?: string; + messageOccurrence: number; +}; + +function readAirForkPoint(meta?: Record | null): AirForkPoint | undefined { + const jetbrains = meta?.["jetbrains"]; + if (!isUnknownRecord(jetbrains)) return undefined; + const air = jetbrains["air"]; + if (!isUnknownRecord(air)) return undefined; + const fork = air["fork"]; + if (!isUnknownRecord(fork) || fork["version"] !== 1) return undefined; + const messageId = fork["messageId"]; + if (typeof messageId !== "string" || messageId.trim().length === 0) { + throw RequestError.invalidParams(undefined, "AIR fork messageId must be a non-empty string"); + } + const messageFingerprint = fork["messageFingerprint"]; + if (messageFingerprint !== undefined + && (typeof messageFingerprint !== "string" || !/^sha256:[0-9a-f]{64}$/.test(messageFingerprint))) { + throw RequestError.invalidParams(undefined, "AIR fork messageFingerprint must be a SHA-256 fingerprint"); + } + const messageOccurrence = fork["messageOccurrence"] ?? 1; + if (!Number.isSafeInteger(messageOccurrence) || (messageOccurrence as number) < 1) { + throw RequestError.invalidParams(undefined, "AIR fork messageOccurrence must be a positive integer"); + } + return { + messageId: messageId.trim(), + ...(typeof messageFingerprint === "string" && {messageFingerprint}), + messageOccurrence: messageOccurrence as number, + }; +} + +function fingerprintAgentMessage(text: string): string { + return `sha256:${createHash("sha256").update(text, "utf8").digest("hex")}`; +} + +function airForkMessageIdCandidates(messageId: string): string[] { + // Older AIR builds sent their visible segment id. Prefer the exact id before its ACP source id. + const visibleSegmentSuffix = /:segment:\d+$/; + const protocolMessageId = messageId.replace(visibleSegmentSuffix, ""); + return protocolMessageId === messageId ? [messageId] : [messageId, protocolMessageId]; +} + +function isUnknownRecord(value: unknown): value is Record { + return typeof value === "object" && value !== null && !Array.isArray(value); +} diff --git a/src/SessionMetadata.ts b/src/SessionMetadata.ts new file mode 100644 index 00000000..50562058 --- /dev/null +++ b/src/SessionMetadata.ts @@ -0,0 +1,17 @@ +import type {ModeKind} from "./app-server/ModeKind"; +import type {ServiceTier} from "./app-server/ServiceTier"; +import type {Model, Thread} from "./app-server/v2"; + +export type SessionMetadata = { + sessionId: string, + currentModelId: string, + models: Model[], + collaborationMode: ModeKind, + modelProvider?: string | null, + currentServiceTier?: ServiceTier | null, + additionalDirectories: string[], +} + +export type SessionMetadataWithThread = SessionMetadata & { + thread: Thread, +} diff --git a/src/__tests__/CodexACPAgent/CodexAcpClient.test.ts b/src/__tests__/CodexACPAgent/CodexAcpClient.test.ts index ba2f8cf5..ced1752a 100644 --- a/src/__tests__/CodexACPAgent/CodexAcpClient.test.ts +++ b/src/__tests__/CodexACPAgent/CodexAcpClient.test.ts @@ -563,6 +563,139 @@ describe('ACP server test', { timeout: 40_000 }, () => { }); }); + it('forks an ACP session through thread/fork with the requested workspace', async () => { + const mockFixture = createCodexMockTestFixture(); + const codexAcpClient = mockFixture.getCodexAcpClient(); + const codexAppServerClient = mockFixture.getCodexAppServerClient(); + + vi.spyOn(codexAppServerClient, "skillsExtraRootsSet").mockResolvedValue(undefined); + vi.spyOn(codexAppServerClient, "listSkills").mockResolvedValue({data: []}); + const threadForkSpy = vi.spyOn(codexAppServerClient, "threadFork").mockResolvedValue({ + thread: {id: "fork-id"} as any, + model: "gpt-5", + modelProvider: "openai", + reasoningEffort: "medium", + serviceTier: null, + } as any); + const threadUnsubscribeSpy = vi.spyOn(codexAppServerClient, "threadUnsubscribe").mockResolvedValue({ + status: "unsubscribed", + }); + vi.spyOn(codexAppServerClient, "listModels").mockResolvedValue({ + data: [createTestModel({id: "gpt-5"})], + nextCursor: null, + }); + + const forked = await codexAcpClient.forkSession({ + sessionId: "source-id", + cwd: "/workspace", + additionalDirectories: ["/workspace/extra"], + mcpServers: [], + }); + + expect(forked.sessionId).toBe("fork-id"); + expect(forked.additionalDirectories).toEqual(["/workspace/extra"]); + expect(threadForkSpy).toHaveBeenCalledWith(expect.objectContaining({ + threadId: "source-id", + cwd: "/workspace", + config: expect.objectContaining({ + projects: { + "/workspace": {trust_level: "trusted"}, + "/workspace/extra": {trust_level: "trusted"}, + }, + }), + })); + expect(threadUnsubscribeSpy).toHaveBeenCalledWith({threadId: "fork-id"}); + }); + + it('maps an AIR fork message id to the containing Codex turn', async () => { + const mockFixture = createCodexMockTestFixture(); + const codexAcpClient = mockFixture.getCodexAcpClient(); + const codexAppServerClient = mockFixture.getCodexAppServerClient(); + + vi.spyOn(codexAppServerClient, "skillsExtraRootsSet").mockResolvedValue(undefined); + vi.spyOn(codexAppServerClient, "listSkills").mockResolvedValue({data: []}); + vi.spyOn(codexAppServerClient, "threadRead").mockResolvedValue({ + thread: { + id: "source-id", + turns: [ + {id: "turn-1", items: [{id: "item-1"}]}, + {id: "turn-2", items: [{id: "agent-message-2"}]}, + ], + }, + } as any); + const threadForkSpy = vi.spyOn(codexAppServerClient, "threadFork").mockResolvedValue({ + thread: {id: "fork-id"}, + model: "gpt-5", + modelProvider: "openai", + reasoningEffort: "medium", + serviceTier: null, + } as any); + vi.spyOn(codexAppServerClient, "listModels").mockResolvedValue({ + data: [createTestModel({id: "gpt-5"})], + nextCursor: null, + }); + + await codexAcpClient.forkSession({ + sessionId: "source-id", + cwd: "/workspace", + _meta: { + jetbrains: {air: {fork: {version: 1, messageId: "agent-message-2:segment:0"}}}, + }, + }); + + expect(threadForkSpy).toHaveBeenCalledWith(expect.objectContaining({ + threadId: "source-id", + lastTurnId: "turn-2", + })); + }); + + it('maps a persisted AIR message fingerprint when Codex item ids changed', async () => { + const mockFixture = createCodexMockTestFixture(); + const codexAcpClient = mockFixture.getCodexAcpClient(); + const codexAppServerClient = mockFixture.getCodexAppServerClient(); + + vi.spyOn(codexAppServerClient, "skillsExtraRootsSet").mockResolvedValue(undefined); + vi.spyOn(codexAppServerClient, "listSkills").mockResolvedValue({data: []}); + vi.spyOn(codexAppServerClient, "threadRead").mockResolvedValue({ + thread: { + id: "source-id", + turns: [ + {id: "turn-1", items: [{type: "agentMessage", id: "new-item-1", text: "Same answer"}]}, + {id: "turn-2", items: [{type: "agentMessage", id: "new-item-2", text: "Same answer"}]}, + ], + }, + } as any); + const threadForkSpy = vi.spyOn(codexAppServerClient, "threadFork").mockResolvedValue({ + thread: {id: "fork-id"}, + model: "gpt-5", + modelProvider: "openai", + reasoningEffort: "medium", + serviceTier: null, + } as any); + vi.spyOn(codexAppServerClient, "listModels").mockResolvedValue({ + data: [createTestModel({id: "gpt-5"})], + nextCursor: null, + }); + + await codexAcpClient.forkSession({ + sessionId: "source-id", + cwd: "/workspace", + _meta: { + jetbrains: {air: {fork: { + version: 1, + messageId: "old-item-2", + messageFingerprint: "sha256:41153d2b46c2869f4021958d44dac18888247fd999507c28970be299a8de4a0f", + messageOccurrence: 2, + }}}, + }, + }); + + expect(threadForkSpy).toHaveBeenCalledWith(expect.objectContaining({ + threadId: "source-id", + lastTurnId: "turn-2", + })); + }); + it('restores collaboration mode for resumed and loaded sessions', async () => { const mockFixture = createCodexMockTestFixture(); const codexAcpAgent = mockFixture.getCodexAcpAgent(); diff --git a/src/__tests__/CodexACPAgent/initialize.test.ts b/src/__tests__/CodexACPAgent/initialize.test.ts index 65350488..9f8fe458 100644 --- a/src/__tests__/CodexACPAgent/initialize.test.ts +++ b/src/__tests__/CodexACPAgent/initialize.test.ts @@ -52,6 +52,7 @@ describe('CodexACPAgent - initialize', () => { list: {}, close: {}, delete: {}, + fork: {}, additionalDirectories: {}, subagents: {}, }, diff --git a/src/__tests__/CodexACPAgent/session-fork.test.ts b/src/__tests__/CodexACPAgent/session-fork.test.ts new file mode 100644 index 00000000..edc4007e --- /dev/null +++ b/src/__tests__/CodexACPAgent/session-fork.test.ts @@ -0,0 +1,39 @@ +import {describe, expect, it, vi} from "vitest"; +import {createCodexMockTestFixture, createTestModel} from "../acp-test-utils"; + +describe("ACP session fork", () => { + it("creates and installs a forked session", async () => { + const fixture = createCodexMockTestFixture(); + const agent = fixture.getCodexAcpAgent(); + const client = fixture.getCodexAcpClient(); + const model = createTestModel({id: "gpt-5"}); + + vi.spyOn(client, "authRequired").mockResolvedValue(false); + vi.spyOn(client, "getAccount").mockResolvedValue({account: null, requiresOpenaiAuth: false}); + vi.spyOn(client, "listSkills").mockResolvedValue({data: []}); + const forkSpy = vi.spyOn(client, "forkSession").mockResolvedValue({ + sessionId: "fork-id", + currentModelId: "gpt-5[medium]", + models: [model], + collaborationMode: "default", + modelProvider: "openai", + currentServiceTier: null, + additionalDirectories: [], + }); + + const response = await agent.forkSession({ + sessionId: "source-id", + cwd: "/workspace", + mcpServers: [], + }); + + expect(response.sessionId).toBe("fork-id"); + expect(agent.getSessionState("fork-id").cwd).toBe("/workspace"); + expect(fixture.getAcpConnectionEvents([])).toEqual([]); + expect(forkSpy).toHaveBeenCalledWith({ + sessionId: "source-id", + cwd: "/workspace", + mcpServers: [], + }); + }); +}); diff --git a/src/index.ts b/src/index.ts index 15ad2256..68df2ccd 100644 --- a/src/index.ts +++ b/src/index.ts @@ -143,6 +143,7 @@ function startAcpServer() { .onRequest(acp.methods.agent.initialize, (ctx) => getAgent().initialize(ctx.params)) .onRequest(acp.methods.agent.session.new, (ctx) => getAgent().newSession(ctx.params)) .onRequest(acp.methods.agent.session.load, (ctx) => getAgent().loadSession(ctx.params)) + .onRequest(acp.methods.agent.session.fork, (ctx) => getAgent().forkSession(ctx.params)) .onRequest(acp.methods.agent.session.list, (ctx) => getAgent().listSessions(ctx.params)) .onRequest(acp.methods.agent.session.delete, (ctx) => getAgent().deleteSession(ctx.params)) .onRequest(acp.methods.agent.session.resume, (ctx) => getAgent().resumeSession(ctx.params))