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
32 changes: 18 additions & 14 deletions src/CodexAcpClient.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -487,6 +490,21 @@ export class CodexAcpClient {
}
}

async forkSession(request: acp.ForkSessionRequest): Promise<SessionMetadata> {
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<SessionMetadataWithThread> {
const additionalDirectories = readAdditionalDirectories(request.cwd, request.additionalDirectories, request._meta);
await this.refreshSkills(request.cwd, additionalDirectories);
Expand Down Expand Up @@ -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) {
Expand Down
59 changes: 46 additions & 13 deletions src/CodexAcpServer.ts
Original file line number Diff line number Diff line change
Expand Up @@ -315,6 +315,7 @@ export class CodexAcpServer {
list: { },
close: { },
delete: { },
fork: { },
additionalDirectories: {},
subagents: {},
};
Expand Down Expand Up @@ -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 ?? [];
Expand All @@ -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;
Expand All @@ -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 = {
Expand Down Expand Up @@ -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),
Expand All @@ -638,16 +648,19 @@ 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,
});
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);
Expand Down Expand Up @@ -733,6 +746,26 @@ export class CodexAcpServer {
};
}

async forkSession(params: acp.ForkSessionRequest): Promise<acp.ForkSessionResponse> {
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<acp.ListSessionsResponse> {
logger.log("Listing sessions...", {cwd: params.cwd, cursor: params.cursor});
await this.checkAuthorization();
Expand Down
134 changes: 134 additions & 0 deletions src/SessionFork.ts
Original file line number Diff line number Diff line change
@@ -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<void>;
createSessionConfig(
cwd: string,
additionalDirectories: string[],
mcpServers: acp.McpServer[],
): Promise<NonNullable<ThreadForkParams["config"]>>;
getResumeModelProvider(): Promise<string>;
fetchAvailableModels(): Promise<Model[]>;
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<SessionMetadata> {
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<string | undefined> {
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<string, unknown> | 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<string, unknown> {
return typeof value === "object" && value !== null && !Array.isArray(value);
}
17 changes: 17 additions & 0 deletions src/SessionMetadata.ts
Original file line number Diff line number Diff line change
@@ -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,
}
Loading