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
52 changes: 45 additions & 7 deletions src/CodexAcpClient.ts
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,15 @@ import * as acp from "@agentclientprotocol/sdk";
import type {CodexAppServerClient} from "./CodexAppServerClient";
import {RequestError} from "@agentclientprotocol/sdk";
import open from "open";
import type {ClientInfo, ServerNotification} from "./app-server";
import type {
ClientInfo,
ServerNotification,
SetDefaultModelParams,
SetDefaultModelResponse
} from "./app-server";
import type {JsonValue} from "./app-server/serde_json/JsonValue";
import type {Model} from "./app-server/v2";
import {ModelId} from "./ModelId";

/**
* API for accessing the Codex App Server using ACP requests.
Expand All @@ -13,12 +20,12 @@ import type {JsonValue} from "./app-server/serde_json/JsonValue";
export class CodexAcpClient {

private readonly codexClient: CodexAppServerClient;
private readonly config: JsonObject | null;
private readonly config: JsonObject;
private readonly modelProvider: string | null;

constructor(codexClient: CodexAppServerClient, codexConfig?: JsonObject, modelProvider?: string) {
this.codexClient = codexClient;
this.config = codexConfig ?? null;
this.config = codexConfig ?? {};
this.modelProvider = modelProvider ?? null;
}

Expand Down Expand Up @@ -71,8 +78,8 @@ export class CodexAcpClient {
/**
* Returns a new session ID.
*/
async newSession(request: acp.NewSessionRequest): Promise<string> {
const response = await this.codexClient.threadStart({
async newSession(request: acp.NewSessionRequest): Promise<SessionMetadata> {
const threadStartResponse = await this.codexClient.threadStart({
config: this.config,
modelProvider: this.modelProvider,
model: null,
Expand All @@ -82,7 +89,16 @@ export class CodexAcpClient {
baseInstructions: null,
developerInstructions: null,
});
return response.thread.id;
const codexModels = await this.fetchAvailableModels();
if (codexModels.length === 0) {
throw new Error("Codex did not return any models");
}
const currentModelId = ModelId.fromThreadResponse(threadStartResponse).toString();
return {
sessionId: threadStartResponse.thread.id,
currentModelId: currentModelId,
models: codexModels
};
}

async sendPrompt(request: acp.PromptRequest, eventHandler: (result: ServerNotification) => void): Promise<void> {
Expand All @@ -106,6 +122,28 @@ export class CodexAcpClient {
await this.codexClient.awaitTurnCompleted();
}

async setModel(params: SetDefaultModelParams): Promise<SetDefaultModelResponse> {
return this.codexClient.setModelRequest(params);
}

private async fetchAvailableModels(): Promise<Model[]> {
const models: Model[] = [];
let cursor: string | null = null;

do {
const response = await this.codexClient.listModels({ cursor, limit: null });
models.push(...response.data);
cursor = response.nextCursor;
} while (cursor);

return models;
}
}

export type JsonObject = { [key in string]?: JsonValue }
export type JsonObject = { [key in string]?: JsonValue }

export type SessionMetadata = {
sessionId: string,
currentModelId: string,
models: Model[]
}
71 changes: 65 additions & 6 deletions src/CodexAcpServer.ts
Original file line number Diff line number Diff line change
@@ -1,12 +1,15 @@
import * as acp from "@agentclientprotocol/sdk";
import {CodexEventHandler} from "./CodexEventHandler";
import {CodexAuthMethods, type CodexAuthRequest} from "./CodexAuthMethod";
import {RequestError} from "@agentclientprotocol/sdk";
import {CodexAcpClient} from "./CodexAcpClient";
import {type ModelInfo, RequestError, type SessionModelState} from "@agentclientprotocol/sdk";
import {CodexAcpClient, type SessionMetadata} from "./CodexAcpClient";
import type {Model} from "./app-server/v2";
import type {ReasoningEffort} from "./app-server";
import {ModelId} from "./ModelId";


export interface SessionState {
sessionId: string,
sessionMetadata: SessionMetadata;
pendingPrompt: AbortController | null;
}

Expand Down Expand Up @@ -52,14 +55,21 @@ export class CodexAcpServer implements acp.Agent {
}
}

const sessionId = await this.codexAcpClient.newSession(_params);
const sessionMetadata = await this.codexAcpClient.newSession(_params);
const {sessionId, currentModelId, models} = sessionMetadata;
this.sessions.set(sessionId, {
sessionId: sessionId,
pendingPrompt: null,
sessionMetadata: sessionMetadata,
pendingPrompt: null
});

const availableModels = this.buildAvailableModels(models);
const sessionModelState: SessionModelState = {
availableModels: availableModels,
currentModelId: currentModelId,
}
return {
sessionId: sessionId,
models: sessionModelState,
};
}

Expand All @@ -80,6 +90,55 @@ export class CodexAcpServer implements acp.Agent {
return {};
}

async setSessionModel(params: acp.SetSessionModelRequest): Promise<acp.SetSessionModelResponse> {
const sessionState = this.sessions.get(params.sessionId);
if (!sessionState) throw new Error(`Session ${params.sessionId} not found`);

const requestedModelId= ModelId.fromString(params.modelId);
const requestedModelName = requestedModelId.model;
const requestedEffort = requestedModelId.effort;

const model = sessionState.sessionMetadata.models.find(m => m.id === requestedModelName);
if (!model) throw new Error(`Unknown model ${params.modelId}`);

const requestedEffortValue = requestedEffort as ReasoningEffort | undefined;
let reasoningEffort: ReasoningEffort;
if (requestedEffortValue) {
const matchedEffort = model.supportedReasoningEfforts.find(
(option) => option.reasoningEffort === requestedEffortValue
)?.reasoningEffort;

if (!matchedEffort) {
throw new Error(`Unsupported reasoning effort ${requestedEffortValue} for model ${requestedModelName}`);
}

reasoningEffort = matchedEffort;
} else {
reasoningEffort = model.defaultReasoningEffort;
}


await this.codexAcpClient.setModel({
model: model.model,
reasoningEffort,
});
sessionState.sessionMetadata.currentModelId = ModelId.fromComponents(model, reasoningEffort).toString();

return {};
}



private buildAvailableModels(models: Model[]): ModelInfo[] {
return models.flatMap((model) =>
model.supportedReasoningEfforts.map((effort) => ({
modelId: ModelId.fromComponents(model, effort.reasoningEffort).toString(),
name: `${model.displayName} (${effort.reasoningEffort})`,
description: `${model.description} ${effort.description}`,
}))
);
}

getSessionState(sessionId: string): SessionState {
const sessionState = this.sessions.get(sessionId);
if (!sessionState) {
Expand Down
17 changes: 13 additions & 4 deletions src/CodexAppServerClient.ts
Original file line number Diff line number Diff line change
@@ -1,15 +1,15 @@
import type {MessageConnection, NotificationMessage} from "vscode-jsonrpc/node";
import type {MessageConnection} from "vscode-jsonrpc/node";
import type {
ClientRequest,
EventMsg,
InitializeParams,
InitializeResponse,
ServerNotification
ServerNotification, SetDefaultModelParams, SetDefaultModelResponse
} from "./app-server";
import type {
AccountLoginCompletedNotification, AccountUpdatedNotification,
GetAccountParams,
GetAccountResponse, LoginAccountParams, LoginAccountResponse, LogoutAccountResponse,
GetAccountResponse, LoginAccountParams, LoginAccountResponse, LogoutAccountResponse, ModelListParams,
ModelListResponse,
ThreadStartParams,
ThreadStartResponse,
TurnCompletedNotification,
Expand Down Expand Up @@ -86,6 +86,15 @@ export class CodexAppServerClient {
});
}


async setModelRequest(params: SetDefaultModelParams): Promise<SetDefaultModelResponse> {
return await this.sendRequest({ method: "setDefaultModel", params });
}

async listModels(params: ModelListParams = {cursor: null, limit: null}): Promise<ModelListResponse> {
return await this.sendRequest({ method: "model/list", params });
}

//TODO support removal (leads to duplicated processing of follow-ups)
onServerNotification(callback: (event: ServerNotification) => void){
this.notificationHandlers.push(callback);
Expand Down
2 changes: 1 addition & 1 deletion src/CodexEventHandler.ts
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@ export class CodexEventHandler {
}

async handleNotification(notification: ServerNotification) {
const session = new ACPSessionConnection(this.connection, this.sessionState.sessionId);
const session = new ACPSessionConnection(this.connection, this.sessionState.sessionMetadata.sessionId);
const updateEvent = await this.createUpdateEvent(notification);
if (updateEvent) {
await session.update(updateEvent);
Expand Down
32 changes: 32 additions & 0 deletions src/ModelId.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,32 @@
import type {ReasoningEffort} from "./app-server";
import type {Model, ThreadStartResponse} from "./app-server/v2";

export class ModelId {
private constructor(
public readonly model: string,
public readonly effort: string | null // TODO: ThreadStartResponse
) {}

static fromComponents(model: Model, effort: ReasoningEffort): ModelId {
return new ModelId(model.id, effort);
}

static fromThreadResponse(response: ThreadStartResponse): ModelId {
return new ModelId(response.model, response.reasoningEffort);
}

static fromString(modelId: string): ModelId {
const parts = modelId.split("/");
const model = parts[0];
const effort = parts[1] ?? null;

if (!model) {
throw new Error(`Invalid modelId format: ${modelId}`);
}
return new ModelId(model, effort);
}

toString(): string {
return this.effort ? `${this.model}/${this.effort}` : this.model;
}
}
14 changes: 11 additions & 3 deletions src/__tests__/CodexACPAgent/CodexAcpClient.test.ts
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
import {describe, expect, it, vi, beforeEach, afterEach} from 'vitest';
import {describe, expect, it, vi, beforeEach} from 'vitest';
import type {CodexAuthRequest} from "../../CodexAuthMethod";
import {createTestFixture, type TestFixture} from "../acp-test-utils";
import type {ServerNotification} from "../../app-server";
Expand All @@ -12,7 +12,7 @@ describe('ACP server test', () => {
vi.clearAllMocks();
});

const ignoredFields = ["thread", "cwd", "id", "createdAt", "path", "threadId", "userAgent", "sandbox", "reasoningEffort"];
const ignoredFields = ["thread", "cwd", "id", "createdAt", "path", "threadId", "userAgent", "sandbox", "reasoningEffort", "conversationId"];

it('should start conversation', async () => {
const codexAcpAgent = fixture.getCodexAcpAgent();
Expand Down Expand Up @@ -79,7 +79,15 @@ describe('ACP server test', () => {

fixture.getCodexAppServerClient().turnStart = vi.fn().mockResolvedValue(undefined);
fixture.getCodexAppServerClient().awaitTurnCompleted = vi.fn().mockResolvedValue(undefined);
fixture.getCodexAcpAgent().getSessionState = vi.fn().mockResolvedValue({ pendingPrompt: null, sessionId: "id" });
const sessionState: SessionState = {
pendingPrompt: null,
sessionMetadata: {
sessionId: "id",
currentModelId: "model-id",
models: [],
}
};
vi.spyOn(codexAcpAgent, "getSessionState").mockReturnValue(sessionState);

await codexAcpAgent.prompt({ sessionId: "id", prompt: [{type: "text", text: ""}] });

Expand Down
Loading