Skip to content
Open
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
44 changes: 27 additions & 17 deletions src/CodexAcpServer.ts
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,7 @@ import {
REASONING_EFFORT_CONFIG_ID,
} from "./ModelConfigOption";
import type {TokenCount} from "./TokenCount";
import {toPromptUsage} from "./TokenCount";
import {subtractTokenCounts, toPromptUsage} from "./TokenCount";
import {CodexCommands} from "./CodexCommands";
import {SteeringQueue} from "./SteeringQueue";
import type {QuotaMeta} from "./QuotaMeta";
Expand Down Expand Up @@ -1878,6 +1878,7 @@ export class CodexAcpServer {
prompt: params.prompt,
});
const sessionState = this.getSessionState(params.sessionId);
const promptStartTotalTokenUsage = sessionState.totalTokenUsage;
sessionState.currentTurnId = null;
sessionState.lastTokenUsage = null;
const activePrompt = this.trackActivePrompt(params.sessionId);
Expand Down Expand Up @@ -1913,7 +1914,7 @@ export class CodexAcpServer {
elicitationHandler);

if (activePrompt.signal.aborted) {
return this.cancelledPromptResponse(sessionState);
return this.cancelledPromptResponse(sessionState, promptStartTotalTokenUsage);
}

const commandPromise = this.availableCommands.tryHandleCommand(params.prompt, sessionState, {
Expand Down Expand Up @@ -1955,14 +1956,14 @@ export class CodexAcpServer {
this.cancelBeforeTurnStarted(activePrompt),
]);
if (commandResult === null) {
return this.cancelledPromptResponse(sessionState);
return this.cancelledPromptResponse(sessionState, promptStartTotalTokenUsage);
}
if (commandResult.handled) {
logger.log("Prompt handled by a command");
await this.codexAcpClient.waitForSessionNotifications(params.sessionId);
if (commandResult.turnCompleted?.turn.status === "interrupted") {
await this.notifyConversationInterrupted(params.sessionId);
return this.cancelledPromptResponse(sessionState);
return this.cancelledPromptResponse(sessionState, promptStartTotalTokenUsage);
}
const error = eventHandler.getFailure();
if (error) {
Expand All @@ -1971,13 +1972,13 @@ export class CodexAcpServer {
}
return {
stopReason: "end_turn",
usage: this.buildPromptUsage(sessionState.lastTokenUsage),
usage: this.buildPromptUsage(sessionState, promptStartTotalTokenUsage),
_meta: this.buildQuotaMeta(sessionState),
};
}

if (this.sessionIsClosing(params.sessionId)) {
return this.cancelledPromptResponse(sessionState);
return this.cancelledPromptResponse(sessionState, promptStartTotalTokenUsage);
}

const modelId = ModelId.fromString(sessionState.currentModelId);
Expand Down Expand Up @@ -2035,14 +2036,14 @@ export class CodexAcpServer {
]);

if (turnCompleted === null) {
return this.cancelledPromptResponse(sessionState);
return this.cancelledPromptResponse(sessionState, promptStartTotalTokenUsage);
}

await this.codexAcpClient.waitForSessionNotifications(params.sessionId);

if (turnCompleted.turn.status === "interrupted") {
await this.notifyConversationInterrupted(params.sessionId);
return this.cancelledPromptResponse(sessionState);
return this.cancelledPromptResponse(sessionState, promptStartTotalTokenUsage);
}

const error = eventHandler.getFailure();
Expand All @@ -2063,7 +2064,7 @@ export class CodexAcpServer {
activePrompt.signal,
);
if (this.promptShouldStop(params.sessionId, activePrompt)) {
return this.cancelledPromptResponse(sessionState);
return this.cancelledPromptResponse(sessionState, promptStartTotalTokenUsage);
}
if (approved && !this.promptShouldStop(params.sessionId, activePrompt)) {
await this.applyCollaborationModeChange(sessionState, DEFAULT_COLLABORATION_MODE);
Expand Down Expand Up @@ -2111,13 +2112,13 @@ export class CodexAcpServer {
]);

if (turnCompleted === null) {
return this.cancelledPromptResponse(sessionState);
return this.cancelledPromptResponse(sessionState, promptStartTotalTokenUsage);
}

await this.codexAcpClient.waitForSessionNotifications(params.sessionId);
if (turnCompleted.turn.status === "interrupted") {
await this.notifyConversationInterrupted(params.sessionId);
return this.cancelledPromptResponse(sessionState);
return this.cancelledPromptResponse(sessionState, promptStartTotalTokenUsage);
}

const implementationError = eventHandler.getFailure();
Expand All @@ -2134,7 +2135,7 @@ export class CodexAcpServer {

return {
stopReason: "end_turn",
usage: this.buildPromptUsage(sessionState.lastTokenUsage),
usage: this.buildPromptUsage(sessionState, promptStartTotalTokenUsage),
_meta: this.buildQuotaMeta(sessionState),
};
} catch (err) {
Expand Down Expand Up @@ -2212,10 +2213,13 @@ export class CodexAcpServer {
}
}

private cancelledPromptResponse(sessionState: SessionState): acp.PromptResponse {
private cancelledPromptResponse(
sessionState: SessionState,
promptStartTotalTokenUsage: TokenCount | null,
): acp.PromptResponse {
return {
stopReason: "cancelled",
usage: this.buildPromptUsage(sessionState.lastTokenUsage),
usage: this.buildPromptUsage(sessionState, promptStartTotalTokenUsage),
_meta: this.buildQuotaMeta(sessionState),
};
}
Expand Down Expand Up @@ -2249,11 +2253,17 @@ export class CodexAcpServer {
};
}

private buildPromptUsage(lastTokenUsage: TokenCount | null): acp.Usage | null {
if (lastTokenUsage == null) {
private buildPromptUsage(
sessionState: SessionState,
promptStartTotalTokenUsage: TokenCount | null,
): acp.Usage | null {
if (sessionState.lastTokenUsage == null || sessionState.totalTokenUsage == null) {
return null;
}
return toPromptUsage(lastTokenUsage);
const promptTokenUsage = promptStartTotalTokenUsage == null
? sessionState.totalTokenUsage
: subtractTokenCounts(sessionState.totalTokenUsage, promptStartTotalTokenUsage);
return toPromptUsage(promptTokenUsage);
}

private async runWithProcessCheck<T>(operation: () => Promise<T>): Promise<T> {
Expand Down
13 changes: 13 additions & 0 deletions src/TokenCount.ts
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,19 @@ export function toTokenCount(usage: TokenUsageBreakdown): TokenCount {
};
}

/**
* Returns the usage accumulated after a cumulative token-count baseline.
*/
export function subtractTokenCounts(total: TokenCount, baseline: TokenCount): TokenCount {
return {
totalTokens: total.totalTokens - baseline.totalTokens,
inputTokens: total.inputTokens - baseline.inputTokens,
cachedInputTokens: total.cachedInputTokens - baseline.cachedInputTokens,
outputTokens: total.outputTokens - baseline.outputTokens,
reasoningOutputTokens: total.reasoningOutputTokens - baseline.reasoningOutputTokens,
};
}

/**
* Maps our per-turn token breakdown to ACP PromptResponse usage fields.
* Cached input tokens are reported as ACP cache reads, and reasoning output
Expand Down
6 changes: 3 additions & 3 deletions src/__tests__/CodexACPAgent/data/token-usage-cancelled.json
Original file line number Diff line number Diff line change
@@ -1,10 +1,10 @@
{
"stopReason": "cancelled",
"usage": {
"totalTokens": 1500,
"inputTokens": 1200,
"totalTokens": 3000,
"inputTokens": 2500,
"cachedReadTokens": 0,
"outputTokens": 300,
"outputTokens": 500,
"thoughtTokens": 0
},
"_meta": {
Expand Down
10 changes: 5 additions & 5 deletions src/__tests__/CodexACPAgent/data/token-usage-end-turn.json
Original file line number Diff line number Diff line change
@@ -1,11 +1,11 @@
{
"stopReason": "end_turn",
"usage": {
"totalTokens": 2500,
"inputTokens": 1500,
"cachedReadTokens": 500,
"outputTokens": 450,
"thoughtTokens": 50
"totalTokens": 5000,
"inputTokens": 3000,
"cachedReadTokens": 1000,
"outputTokens": 900,
"thoughtTokens": 100
},
"_meta": {
"quota": {
Expand Down
Original file line number Diff line number Diff line change
@@ -1,10 +1,10 @@
{
"stopReason": "end_turn",
"usage": {
"totalTokens": 1500,
"inputTokens": 700,
"totalTokens": 3500,
"inputTokens": 2300,
"cachedReadTokens": 500,
"outputTokens": 200,
"outputTokens": 600,
"thoughtTokens": 100
},
"_meta": {
Expand Down
33 changes: 33 additions & 0 deletions src/__tests__/CodexACPAgent/data/token-usage-prompt-delta.json
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
{
"stopReason": "end_turn",
"usage": {
"totalTokens": 3200,
"inputTokens": 1600,
"cachedReadTokens": 1000,
"outputTokens": 600,
"thoughtTokens": 100
},
"_meta": {
"quota": {
"token_count": {
"totalTokens": 1500,
"inputTokens": 700,
"cachedInputTokens": 500,
"outputTokens": 300,
"reasoningOutputTokens": 100
},
"model_usage": [
{
"model": "model-id",
"token_count": {
"totalTokens": 1500,
"inputTokens": 700,
"cachedInputTokens": 500,
"outputTokens": 300,
"reasoningOutputTokens": 100
}
}
]
}
}
}
58 changes: 55 additions & 3 deletions src/__tests__/CodexACPAgent/token-usage-events.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ import { describe, it, expect, vi, beforeEach } from 'vitest';
import type { ServerNotification } from '../../app-server';
import { createCodexMockTestFixture, createTestSessionState, type CodexMockTestFixture } from '../acp-test-utils';
import type { TokenUsageBreakdown } from '../../app-server/v2';
import type { TokenCount } from '../../TokenCount';

function createTokenUsageNotification(
sessionId: string,
Expand Down Expand Up @@ -30,7 +31,11 @@ describe('Token Usage Events', () => {
vi.clearAllMocks();
});
describe('PromptResponse usage', () => {
function setupPromptWithTokenUsage(notifications: ServerNotification[], turnStatus: string = "completed") {
function setupPromptWithTokenUsage(
notifications: ServerNotification[],
turnStatus: string = "completed",
initialTotalTokenUsage: TokenCount | null = null,
) {
const codexAcpAgent = mockFixture.getCodexAcpAgent();

mockFixture.getCodexAppServerClient().turnStart = vi.fn().mockResolvedValue({
Expand All @@ -49,7 +54,10 @@ describe('Token Usage Events', () => {
};
});

vi.spyOn(codexAcpAgent, 'getSessionState').mockReturnValue(createTestSessionState({ sessionId }));
vi.spyOn(codexAcpAgent, 'getSessionState').mockReturnValue(createTestSessionState({
sessionId,
totalTokenUsage: initialTotalTokenUsage,
}));

return codexAcpAgent;
}
Expand Down Expand Up @@ -133,7 +141,7 @@ describe('Token Usage Events', () => {
);
});

it('should use last token usage from multiple updates', async () => {
it('should include cumulative token usage from multiple updates', async () => {
const notifications: ServerNotification[] = [
createTokenUsageNotification(sessionId, {
total: { totalTokens: 1000, inputTokens: 800, cachedInputTokens: 0, cacheWriteInputTokens: 0, outputTokens: 200, reasoningOutputTokens: 0 },
Expand Down Expand Up @@ -163,6 +171,50 @@ describe('Token Usage Events', () => {
'data/token-usage-multiple-updates.json'
);
});

it('should report usage accumulated since the previous prompt', async () => {
const initialTotalTokenUsage: TokenCount = {
totalTokens: 2000,
inputTokens: 1600,
cachedInputTokens: 0,
outputTokens: 400,
reasoningOutputTokens: 0,
};
const tokenUsageNotification = createTokenUsageNotification(sessionId, {
total: {
totalTokens: 5200,
inputTokens: 4200,
cachedInputTokens: 1000,
cacheWriteInputTokens: 0,
outputTokens: 1000,
reasoningOutputTokens: 100,
},
last: {
totalTokens: 1500,
inputTokens: 1200,
cachedInputTokens: 500,
cacheWriteInputTokens: 0,
outputTokens: 300,
reasoningOutputTokens: 100,
},
modelContextWindow: 128000,
});

const codexAcpAgent = setupPromptWithTokenUsage(
[tokenUsageNotification],
"completed",
initialTotalTokenUsage,
);

const response = await codexAcpAgent.prompt({
sessionId,
prompt: [{ type: 'text', text: 'test prompt' }],
});

await expect(`${JSON.stringify(response, null, 2)}\n`).toMatchFileSnapshot(
'data/token-usage-prompt-delta.json'
);
});
});

describe('session/update usage_update', () => {
Expand Down