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
8 changes: 8 additions & 0 deletions src/CodexAcpClient.ts
Original file line number Diff line number Diff line change
Expand Up @@ -217,6 +217,10 @@ export class CodexAcpClient {
return response.requiresOpenaiAuth && !response.account;
}

hasGatewayAuth(): boolean {
return this.gatewayConfig !== null;
}

async getAccount(): Promise<GetAccountResponse> {
return this.codexClient.accountRead({refreshToken: false});
}
Expand All @@ -238,6 +242,7 @@ export class CodexAcpClient {
sessionId: request.sessionId,
currentModelId: currentModelId,
models: codexModels,
modelProvider: response.modelProvider,
currentServiceTier: response.serviceTier as ServiceTier ?? null,
additionalDirectories,
}
Expand All @@ -264,6 +269,7 @@ export class CodexAcpClient {
sessionId: request.sessionId,
currentModelId: currentModelId,
models: codexModels,
modelProvider: response.modelProvider,
currentServiceTier: response.serviceTier as ServiceTier ?? null,
thread: historyResponse.thread,
additionalDirectories,
Expand All @@ -289,6 +295,7 @@ export class CodexAcpClient {
sessionId: response.thread.id,
currentModelId: currentModelId,
models: codexModels,
modelProvider: response.modelProvider,
currentServiceTier: response.serviceTier as ServiceTier ?? null,
additionalDirectories,
};
Expand Down Expand Up @@ -698,6 +705,7 @@ export type SessionMetadata = {
sessionId: string,
currentModelId: string,
models: Model[],
modelProvider?: string | null,
currentServiceTier?: ServiceTier | null,
additionalDirectories: string[],
}
Expand Down
79 changes: 67 additions & 12 deletions src/CodexAcpServer.ts
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@ import {RequestError, type SessionId, type SessionModeState} from "@agentclientp
import {CodexEventHandler} from "./CodexEventHandler";
import {CodexApprovalHandler} from "./CodexApprovalHandler";
import {CodexElicitationHandler} from "./CodexElicitationHandler";
import {type CodexAuthRequest, getCodexAuthMethods} from "./CodexAuthMethod";
import {type CodexAuthRequest, getCodexAuthMethods, isCodexAuthRequest} from "./CodexAuthMethod";
import {CodexAcpClient, type SessionMetadata, type SessionMetadataWithThread} from "./CodexAcpClient";
import type {McpStartupResult} from "./CodexAppServerClient";
import {ACPSessionConnection, type AcpClientConnection, type UpdateSessionEvent} from "./ACPSessionConnection";
Expand Down Expand Up @@ -79,6 +79,8 @@ export interface SessionState {
modelContextWindow: number | null;
rateLimits: RateLimitsMap | null;
account: Account | null;
authConfigured: boolean;
authProvider: string | null;
cwd: string;
additionalDirectories: string[];
fastModeEnabled: boolean;
Expand All @@ -87,6 +89,11 @@ export interface SessionState {
terminalOutputMode: TerminalOutputMode;
}

interface ActiveAuthState {
account: Account | null;
authConfigured: boolean;
}

interface PendingMcpStartupSession {
requestedServers: Set<string>;
afterVersion: number;
Expand Down Expand Up @@ -156,7 +163,8 @@ export class CodexAcpServer {
this.availableCommands = new CodexCommands(
connection,
codexAcpClient,
(operation) => this.runWithProcessCheck(operation)
(operation) => this.runWithProcessCheck(operation),
() => this.refreshSessionsAuthState(null)
);
}

Expand Down Expand Up @@ -247,6 +255,7 @@ export class CodexAcpServer {
async handleError(e: Error){
if (e.message.includes("log out") || e.message.includes("cloud requirements")) {
await this.runWithProcessCheck(() => this.codexAcpClient.logout());
await this.refreshSessionsAuthState(null);
throw RequestError.internalError(`${(e.message)}\n\nYou have been logged out. Please try again.`);
}
}
Expand Down Expand Up @@ -346,9 +355,10 @@ export class CodexAcpServer {
}

const {sessionId, currentModelId, models} = sessionMetadata;
let account: Account | null;
const authProvider = sessionMetadata.modelProvider ?? this.codexAcpClient.getModelProvider();
let authState: ActiveAuthState;
try {
account = await this.getActiveAccount();
authState = await this.getAuthStateForProvider(authProvider);
} catch (err) {
if (resumeSubscribed && requestedSessionGeneration !== null) {
await this.cleanupStaleSessionOpen(sessionId, requestedSessionGeneration);
Expand All @@ -375,7 +385,9 @@ export class CodexAcpServer {
totalTokenUsage: null,
modelContextWindow: null,
rateLimits: null,
account: account,
account: authState.account,
authConfigured: authState.authConfigured,
authProvider: authProvider,
cwd: request.cwd,
additionalDirectories: sessionMetadata.additionalDirectories,
fastModeEnabled: sessionMetadata.currentServiceTier === "fast",
Expand All @@ -401,12 +413,36 @@ export class CodexAcpServer {
return [sessionId, sessionModelState, sessionModeState];
}

private async getActiveAccount(){
if (this.codexAcpClient.getModelProvider()) {
return null
private async getAuthStateForProvider(authProvider: string | null): Promise<ActiveAuthState> {
if (!this.authProviderUsesOpenAiAccount(authProvider)) {
return {
account: null,
authConfigured: true,
};
}
const accountResponse = await this.runWithProcessCheck(() => this.codexAcpClient.getAccount());
return accountResponse.account;
return {
account: accountResponse.account,
authConfigured: accountResponse.account !== null || !accountResponse.requiresOpenaiAuth,
};
}

private authProviderUsesOpenAiAccount(authProvider: string | null): boolean {
return authProvider === null || authProvider === "openai";
}

private authProvidersMatch(a: string | null, b: string | null): boolean {
if (this.authProviderUsesOpenAiAccount(a) && this.authProviderUsesOpenAiAccount(b)) {
return true;
}
return a === b;
}

private getAuthProviderForAuthenticateRequest(request: acp.AuthenticateRequest): string | null {
if (isCodexAuthRequest(request) && request.methodId === "gateway") {
return "custom-gateway";
}
return null;
}

async loadSession(params: acp.LoadSessionRequest): Promise<LegacyLoadSessionResponse> {
Expand Down Expand Up @@ -565,16 +601,32 @@ export class CodexAcpServer {
logger.log("Authenticate request failed");
throw RequestError.invalidParams();
}
await this.refreshSessionsAuthState(this.getAuthProviderForAuthenticateRequest(_params));
logger.log("Authenticate request completed");
return { };
}

async logout(_params: acp.LogoutRequest): Promise<void> {
logger.log("Logout request received");
await this.runWithProcessCheck(() => this.codexAcpClient.logout());
await this.refreshSessionsAuthState(null);
logger.log("Logout request completed");
}

private async refreshSessionsAuthState(authProvider: string | null): Promise<void> {
if (this.sessions.size === 0) return;

const sessionsToRefresh = [...this.sessions.values()]
.filter(sessionState => this.authProvidersMatch(sessionState.authProvider, authProvider));
if (sessionsToRefresh.length === 0) return;

const authState = await this.getAuthStateForProvider(authProvider);
for (const sessionState of sessionsToRefresh) {
sessionState.account = authState.account;
sessionState.authConfigured = authState.authConfigured;
}
}

async setSessionMode(
_params: acp.SetSessionModeRequest,
): Promise<acp.SetSessionModeResponse> {
Expand Down Expand Up @@ -798,9 +850,10 @@ export class CodexAcpServer {
}

const {sessionId, currentModelId, models, thread} = sessionMetadata;
let account: Account | null;
const authProvider = sessionMetadata.modelProvider ?? this.codexAcpClient.getModelProvider();
let authState: ActiveAuthState;
try {
account = await this.getActiveAccount();
authState = await this.getAuthStateForProvider(authProvider);
} catch (err) {
if (subscribed) {
await this.cleanupStaleSessionOpen(request.sessionId, requestedSessionGeneration);
Expand All @@ -826,7 +879,9 @@ export class CodexAcpServer {
totalTokenUsage: null,
modelContextWindow: null,
rateLimits: null,
account: account,
account: authState.account,
authConfigured: authState.authConfigured,
authProvider: authProvider,
cwd: request.cwd,
additionalDirectories: sessionMetadata.additionalDirectories,
fastModeEnabled: sessionMetadata.currentServiceTier === "fast",
Expand Down
8 changes: 7 additions & 1 deletion src/CodexCommands.ts
Original file line number Diff line number Diff line change
Expand Up @@ -21,19 +21,24 @@ export type CommandHandleOptions = {
onTurnStarted?: (turnId: string, threadId: string) => void;
};

export type LogoutHandler = () => void | Promise<void>;

export class CodexCommands {
private readonly connection: AcpClientConnection;
private readonly codexAcpClient: CodexAcpClient;
private readonly runWithProcessCheck: <T>(operation: () => Promise<T>) => Promise<T>;
private readonly onLogout: LogoutHandler;

constructor(
connection: AcpClientConnection,
codexAcpClient: CodexAcpClient,
runWithProcessCheck: <T>(operation: () => Promise<T>) => Promise<T>
runWithProcessCheck: <T>(operation: () => Promise<T>) => Promise<T>,
onLogout: LogoutHandler = () => {}
) {
this.connection = connection;
this.codexAcpClient = codexAcpClient;
this.runWithProcessCheck = runWithProcessCheck;
this.onLogout = onLogout;
}

async publish(sessionId: string): Promise<void> {
Expand Down Expand Up @@ -198,6 +203,7 @@ export class CodexCommands {
}
case "logout": {
await this.runWithProcessCheck(() => this.codexAcpClient.logout());
await this.onLogout();
const session = new ACPSessionConnection(this.connection, sessionId);
await session.update({
sessionUpdate: "agent_message_chunk",
Expand Down
37 changes: 34 additions & 3 deletions src/CodexEventHandler.ts
Original file line number Diff line number Diff line change
Expand Up @@ -589,9 +589,15 @@ export class CodexEventHandler {
}

private async createErrorEvent(params: ErrorNotification): Promise<UpdateSessionEvent> {
const error = params.error.codexErrorInfo
if (error == "unauthorized" || error == "usageLimitExceeded" || this.getHttpStatusCode(error) == 401) {
this.failure = RequestError.authRequired();
const error = params.error.codexErrorInfo;
if (error === "usageLimitExceeded") {
this.failure = RequestError.internalError(
this.createTurnErrorData(params.error),
);
} else if (this.isAuthenticationRequiredError(error)) {
this.failure = this.sessionState.authConfigured
? RequestError.internalError(this.createTurnErrorData(params.error))
: RequestError.authRequired(this.createTurnErrorData(params.error), params.error.message);
}
return {
sessionUpdate: "agent_message_chunk",
Expand All @@ -602,6 +608,10 @@ export class CodexEventHandler {
}
}

private isAuthenticationRequiredError(error: CodexErrorInfo | null): boolean {
return error === "unauthorized" || this.getHttpStatusCode(error) === 401;
}

private getHttpStatusCode(error: CodexErrorInfo | null): number | null {
if (error !== null && typeof error === "object") {
if ("httpConnectionFailed" in error) {
Expand All @@ -617,6 +627,27 @@ export class CodexEventHandler {
return null;
}

private createTurnErrorData(error: ErrorNotification["error"]): {
message: string;
codexErrorInfo?: CodexErrorInfo;
additionalDetails?: string;
} {
const data: {
message: string;
codexErrorInfo?: CodexErrorInfo;
additionalDetails?: string;
} = {
message: error.additionalDetails ?? error.message,
};
if (error.codexErrorInfo !== null) {
data.codexErrorInfo = error.codexErrorInfo;
}
if (error.additionalDetails !== null) {
data.additionalDetails = error.additionalDetails;
}
return data;
}

private handleTokenUsageUpdated(params: ThreadTokenUsageUpdatedNotification): void {
this.sessionState.lastTokenUsage = toTokenCount(params.tokenUsage.last);
this.sessionState.totalTokenUsage = toTokenCount(params.tokenUsage.total);
Expand Down
Loading