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
26 changes: 26 additions & 0 deletions src/CodexAcpServer.ts
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,8 @@ import {
createUserMessageChunk,
} from "./ContentChunks";
import {sameThreadGoalSnapshot, type ThreadGoalSnapshot, toThreadGoalSnapshot,} from "./ThreadGoalSnapshot";
import {OpenAiPricingProvider, type PricingProvider} from "./PricingProvider";
import {SessionCostTracker} from "./SessionCostTracker";
import {randomUUID} from "node:crypto";
import {once} from "node:events";
import {
Expand Down Expand Up @@ -144,6 +146,7 @@ export interface SessionState {
sessionTitle: string | null;
sessionTitleSource: "unset" | "fallback" | "explicit" | "unknown";
sessionFailure?: SessionFailure;
costTracker: SessionCostTracker;
}

export type SessionFailureCategory =
Expand Down Expand Up @@ -240,6 +243,7 @@ export class CodexAcpServer {
private readonly getExitCode: () => number | null;
private readonly getRecentStderr: () => string;
private readonly sessionFailureEpoch: string;
private readonly pricingProvider: PricingProvider;
private availableCommands: CodexCommands;
private clientInfo: acp.Implementation | null;
private clientCapabilities: acp.ClientCapabilities | null;
Expand All @@ -266,6 +270,7 @@ export class CodexAcpServer {
getExitCode?: () => number | null,
getRecentStderr?: () => string,
codexProcessState?: CodexProcessState,
pricingProvider: PricingProvider = new OpenAiPricingProvider(),
) {
this.sessions = new Map();
this.pendingMcpStartupSessions = new Map();
Expand All @@ -284,6 +289,7 @@ export class CodexAcpServer {
this.getExitCode = getExitCode ?? (() => this.codexProcessState?.connection.process.exitCode ?? null);
this.getRecentStderr = getRecentStderr ?? (() => this.codexProcessState?.stderr ?? "");
this.sessionFailureEpoch = randomUUID();
this.pricingProvider = pricingProvider;
this.clientInfo = null;
this.clientCapabilities = null;
this.terminalOutputMode = "terminal_output_delta";
Expand Down Expand Up @@ -600,6 +606,11 @@ export class CodexAcpServer {
const sessionMcpServers = this.resolveSessionMcpServers(requestedMcpServers, "sessionId" in request);
const currentModel = this.findCurrentModel(models, currentModelId);
const currentModelSupportsFast = modelSupportsFast(currentModel);
const costTracker = await this.createSessionCostTracker(
models,
authProvider,
"sessionId" in request,
);
const sessionState: SessionState = {
sessionId: sessionId,
currentModelId: currentModelId,
Expand All @@ -626,6 +637,7 @@ export class CodexAcpServer {
goalRevision: 0,
sessionTitle: null,
sessionTitleSource: "sessionId" in request ? "unknown" : "unset",
costTracker,
};
this.sessions.set(sessionId, sessionState);
resumeSubscribed = false;
Expand Down Expand Up @@ -666,6 +678,18 @@ export class CodexAcpServer {
return authProvider === null || authProvider === "openai";
}

private async createSessionCostTracker(
models: readonly Model[],
authProvider: string | null,
baselineInitialUsage: boolean,
): Promise<SessionCostTracker> {
if (!this.authProviderUsesOpenAiAccount(authProvider)) {
return SessionCostTracker.disabled();
}
const pricing = await this.pricingProvider.getPricing(models);
return new SessionCostTracker(pricing, baselineInitialUsage);
}

private authProvidersMatch(a: string | null, b: string | null): boolean {
if (this.authProviderUsesOpenAiAccount(a) && this.authProviderUsesOpenAiAccount(b)) {
return true;
Expand Down Expand Up @@ -1598,6 +1622,7 @@ export class CodexAcpServer {
const sessionMcpServers = this.resolveSessionMcpServers(requestedMcpServers, true);
const currentModel = this.findCurrentModel(models, currentModelId);
const currentModelSupportsFast = modelSupportsFast(currentModel);
const costTracker = await this.createSessionCostTracker(models, authProvider, true);
const sessionState: SessionState = {
sessionId: sessionId,
currentModelId: currentModelId,
Expand All @@ -1624,6 +1649,7 @@ export class CodexAcpServer {
goalRevision: 0,
sessionTitle: null,
sessionTitleSource: "unset",
costTracker,
};
this.sessions.set(sessionId, sessionState);
subscribed = false;
Expand Down
10 changes: 10 additions & 0 deletions src/CodexEventHandler.ts
Original file line number Diff line number Diff line change
Expand Up @@ -1233,10 +1233,20 @@ export class CodexEventHandler {
return null;
}

const cost = this.sessionState.totalTokenUsage === null || this.sessionState.lastTokenUsage === null
? null
: this.sessionState.costTracker.update(
this.sessionState.totalTokenUsage,
this.sessionState.lastTokenUsage,
this.sessionState.currentModelId,
this.sessionState.fastModeEnabled,
);

return {
sessionUpdate: "usage_update",
used,
size,
...(cost === null ? {} : {cost}),
};
}

Expand Down
144 changes: 144 additions & 0 deletions src/PricingProvider.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,144 @@
import type {Model} from "./app-server/v2";
import {logger} from "./Logger";

export interface TokenRates {
input: number;
cachedInput: number;
output: number;
}

export interface TierPricing {
shortContext: TokenRates;
longContext?: TokenRates;
}

export interface ModelPricing {
standard: TierPricing;
fast?: TierPricing;
}

export type ModelPricingSnapshot = ReadonlyMap<string, ModelPricing>;

export interface PricingProvider {
getPricing(models: readonly Model[]): Promise<ModelPricingSnapshot>;
}

const OPENAI_PRICING_URL = "https://developers.openai.com/api/docs/pricing.md";
const PRICING_FETCH_TIMEOUT_MS = 5_000;

export class OpenAiPricingProvider implements PricingProvider {
private pricingDocument: Promise<string> | null = null;

async getPricing(models: readonly Model[]): Promise<ModelPricingSnapshot> {
if (models.length === 0) return new Map();

try {
const markdown = await this.getPricingDocument();
return parseModelPricing(markdown, models.map(model => model.id));
} catch (error) {
logger.error("Failed to load OpenAI model pricing", error);
return new Map();
}
}

private async getPricingDocument(): Promise<string> {
if (this.pricingDocument === null) {
this.pricingDocument = this.fetchPricingDocument().catch(error => {
this.pricingDocument = null;
throw error;
});
}
return await this.pricingDocument;
}

private async fetchPricingDocument(): Promise<string> {
const response = await fetch(OPENAI_PRICING_URL, {
headers: {accept: "text/markdown"},
signal: AbortSignal.timeout(PRICING_FETCH_TIMEOUT_MS),
});
if (!response.ok) {
throw new Error(`OpenAI pricing request failed with HTTP ${response.status}`);
}
return await response.text();
}
}

function parseModelPricing(markdown: string, modelIds: readonly string[]): ModelPricingSnapshot {
const requestedModels = new Map<string, string[]>();
for (const modelId of modelIds) {
const pricingId = pricingModelId(modelId);
const matchingIds = requestedModels.get(pricingId) ?? [];
matchingIds.push(modelId);
requestedModels.set(pricingId, matchingIds);
}

const standard = parsePricingSection(markdown, "Standard pricing data", requestedModels);
const fast = parsePricingSection(markdown, "Fast pricing data", requestedModels);
const result = new Map<string, ModelPricing>();
for (const [modelId, standardPricing] of standard) {
const fastPricing = fast.get(modelId);
result.set(modelId, {
standard: standardPricing,
...(fastPricing === undefined ? {} : {fast: fastPricing}),
});
}
return result;
}

function parsePricingSection(
markdown: string,
heading: string,
requestedModels: ReadonlyMap<string, readonly string[]>,
): Map<string, TierPricing> {
const headingText = `### ${heading}`;
const start = markdown.indexOf(headingText);
if (start < 0) return new Map();
const nextHeading = markdown.indexOf("\n### ", start + headingText.length);
const section = markdown.slice(start, nextHeading < 0 ? undefined : nextHeading);
const result = new Map<string, TierPricing>();

for (const line of section.split(/\r?\n/)) {
if (!line.startsWith("|")) continue;
const cells = line.split("|").slice(1, -1).map(cell => cell.trim());
if (cells.length < 9 || cells[0] === "Model" || cells[0]?.startsWith("---")) continue;

const documentedModel = cells[0]?.replace(/\s+\([^)]*\)\s*$/, "");
if (documentedModel === undefined) continue;
const matchingModelIds = requestedModels.get(documentedModel);
if (matchingModelIds === undefined) continue;

const shortContext = parseTokenRates(cells[1], cells[2], cells[4]);
if (shortContext === null) continue;
const longContext = parseTokenRates(cells[5], cells[6], cells[8]);
const pricing: TierPricing = {
shortContext,
...(longContext === null ? {} : {longContext}),
};
for (const modelId of matchingModelIds) {
result.set(modelId, pricing);
}
}
return result;
}

function parseTokenRates(
inputValue: string | undefined,
cachedInputValue: string | undefined,
outputValue: string | undefined,
): TokenRates | null {
const input = parseUsdRate(inputValue);
const cachedInput = parseUsdRate(cachedInputValue);
const output = parseUsdRate(outputValue);
if (input === null || cachedInput === null || output === null) return null;
return {input, cachedInput, output};
}

function parseUsdRate(value: string | undefined): number | null {
if (value === undefined || value === "-") return null;
const amount = Number(value.replace(/[$,]/g, ""));
return Number.isFinite(amount) && amount >= 0 ? amount : null;
}

function pricingModelId(modelId: string): string {
return modelId.toLowerCase().replace(/-\d{4}-\d{2}-\d{2}$/, "");
}
93 changes: 93 additions & 0 deletions src/SessionCostTracker.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,93 @@
import type {Cost} from "@agentclientprotocol/sdk";
import type {TokenCount} from "./TokenCount";
import type {ModelPricingSnapshot, TierPricing, TokenRates} from "./PricingProvider";

const LONG_CONTEXT_THRESHOLD = 272_000;
const TOKENS_PER_MILLION = 1_000_000;

export class SessionCostTracker {
private previousTotalUsage: TokenCount | null = null;
private amountUsd = 0;
private available: boolean;

constructor(
private readonly pricing: ModelPricingSnapshot,
private baselineInitialUsage = false,
) {
this.available = pricing.size > 0;
}

static disabled(): SessionCostTracker {
return new SessionCostTracker(new Map());
}

update(
totalUsage: TokenCount,
lastUsage: TokenCount,
currentModelId: string,
fastModeEnabled: boolean,
): Cost | null {
const delta = usageDelta(totalUsage, this.previousTotalUsage);
this.previousTotalUsage = {...totalUsage};
if (!this.available || delta === null) {
this.available = false;
return null;
}

const modelId = currentModelId.replace(/\[[^\]]*]$/, "");
const modelPricing = this.pricing.get(modelId);
const tierPricing = fastModeEnabled ? modelPricing?.fast : modelPricing?.standard;
const rates = selectRates(tierPricing, lastUsage);
if (rates === null) {
this.available = false;
return null;
}

if (this.baselineInitialUsage) {
this.baselineInitialUsage = false;
return {amount: this.amountUsd, currency: "USD"};
}

const incrementalCost = (
delta.inputTokens * rates.input
+ delta.cachedInputTokens * rates.cachedInput
+ delta.outputTokens * rates.output
) / TOKENS_PER_MILLION;
if (!Number.isFinite(incrementalCost) || incrementalCost < 0) {
this.available = false;
return null;
}

this.amountUsd += incrementalCost;
return {amount: this.amountUsd, currency: "USD"};
}
}

function selectRates(pricing: TierPricing | undefined, lastUsage: TokenCount): TokenRates | null {
if (pricing === undefined) return null;
const inputTokens = lastUsage.inputTokens + lastUsage.cachedInputTokens;
if (inputTokens > LONG_CONTEXT_THRESHOLD) {
return pricing.longContext ?? null;
}
return pricing.shortContext;
}

function usageDelta(current: TokenCount, previous: TokenCount | null): TokenCount | null {
if (previous === null) return {...current};
if (
current.totalTokens < previous.totalTokens
|| current.inputTokens < previous.inputTokens
|| current.cachedInputTokens < previous.cachedInputTokens
|| current.outputTokens < previous.outputTokens
|| current.reasoningOutputTokens < previous.reasoningOutputTokens
) {
return null;
}
return {
totalTokens: current.totalTokens - previous.totalTokens,
inputTokens: current.inputTokens - previous.inputTokens,
cachedInputTokens: current.cachedInputTokens - previous.cachedInputTokens,
outputTokens: current.outputTokens - previous.outputTokens,
reasoningOutputTokens: current.reasoningOutputTokens - previous.reasoningOutputTokens,
};
}
Loading