From af639257d751f2cef543b588a89932b99fa988a5 Mon Sep 17 00:00:00 2001 From: "seth.pan" Date: Thu, 3 Sep 2026 23:52:05 +0800 Subject: [PATCH] fix: propagate CodeWhisperer cache token usage instead of hardcoding zero The SDK streaming and non-streaming response paths emitted cache_creation_input_tokens and cache_read_input_tokens as hardcoded zero literals, even though the SDK's TokenUsage type carries optional cacheReadInputTokens and cacheWriteInputTokens fields (see @aws/codewhisperer-streaming-client dist-types/models/models_0.d.ts). Capture the cache fields from MetadataEvent.tokenUsage when present in transformSdkStream and ResponseHandler.handleSdkNonStreaming, and propagate them through convertToOpenAI's message_delta usage emission. The non-SDK paths (raw HTTP event stream) cannot report cache tokens because that wire format only exposes contextUsagePercentage; those sites are left at 0 with a comment explaining why. Pure observability: no request payload is changed, no cachePoint is added, no billing behavior is modified. --- src/__tests__/cache-token-usage.test.ts | 185 ++++++++++++++++++ src/core/request/response-handler.ts | 25 ++- src/plugin/streaming/openai-converter.ts | 4 +- .../streaming/sdk-stream-transformer.ts | 15 +- src/plugin/streaming/stream-transformer.ts | 4 + 5 files changed, 226 insertions(+), 7 deletions(-) create mode 100644 src/__tests__/cache-token-usage.test.ts diff --git a/src/__tests__/cache-token-usage.test.ts b/src/__tests__/cache-token-usage.test.ts new file mode 100644 index 0000000..2abe028 --- /dev/null +++ b/src/__tests__/cache-token-usage.test.ts @@ -0,0 +1,185 @@ +import { describe, expect, test } from 'bun:test' +import { ResponseHandler } from '../core/request/response-handler' +import { transformSdkStream } from '../plugin/streaming/sdk-stream-transformer.js' +import { transformKiroStream } from '../plugin/streaming/stream-transformer.js' + +const MODEL = 'claude-opus-5' + +function sdkStreamOf(events: any[]) { + return { + generateAssistantResponseResponse: (async function* () { + for (const event of events) yield event + })() + } +} + +async function lastUsageChunk(events: any[]): Promise { + const chunks: any[] = [] + for await (const chunk of transformSdkStream(sdkStreamOf(events), MODEL, 'conversation-1')) { + chunks.push(chunk) + } + for (let i = chunks.length - 1; i >= 0; i--) { + if (chunks[i]?.usage) return chunks[i].usage + } + return null +} + +describe('cache token usage — SDK streaming', () => { + test('emits non-zero cache_creation_input_tokens from metadataEvent.tokenUsage', async () => { + const usage = await lastUsageChunk([ + { assistantResponseEvent: { content: 'Answer.' } }, + { + metadataEvent: { + tokenUsage: { + inputTokens: 100, + outputTokens: 25, + cacheReadInputTokens: 800, + cacheWriteInputTokens: 1200 + } + } + } + ]) + + expect(usage).not.toBeNull() + expect(usage.cache_read_input_tokens).toBe(800) + expect(usage.cache_creation_input_tokens).toBe(1200) + }) + + test('emits zeros when metadataEvent has no tokenUsage at all', async () => { + const usage = await lastUsageChunk([ + { assistantResponseEvent: { content: 'Answer.' } }, + { metadataEvent: {} } + ]) + + expect(usage).not.toBeNull() + expect(usage.cache_read_input_tokens).toBe(0) + expect(usage.cache_creation_input_tokens).toBe(0) + }) + + test('emits zeros when tokenUsage omits cache fields (SDK type defines them as optional)', async () => { + const usage = await lastUsageChunk([ + { assistantResponseEvent: { content: 'Answer.' } }, + { + metadataEvent: { + tokenUsage: { + inputTokens: 50, + outputTokens: 10 + // cacheReadInputTokens / cacheWriteInputTokens intentionally absent + } + } + } + ]) + + expect(usage).not.toBeNull() + expect(usage.cache_read_input_tokens).toBe(0) + expect(usage.cache_creation_input_tokens).toBe(0) + }) + + test('does not crash when no metadataEvent ever arrives', async () => { + const usage = await lastUsageChunk([{ assistantResponseEvent: { content: 'Answer.' } }]) + + expect(usage).not.toBeNull() + expect(usage.cache_read_input_tokens).toBe(0) + expect(usage.cache_creation_input_tokens).toBe(0) + }) +}) + +describe('cache token usage — SDK non-streaming', () => { + const handler = new ResponseHandler() + + async function readJsonUsage(sdkResponse: any): Promise { + const response = await handler.handleSdkSuccess(sdkResponse, MODEL, 'conversation-1', false) + const body = await response.json() + return body.usage + } + + test('emits non-zero cache fields when metadataEvent carries them', async () => { + const usage = await readJsonUsage( + sdkStreamOf([ + { assistantResponseEvent: { content: 'Answer.' } }, + { + metadataEvent: { + tokenUsage: { + inputTokens: 200, + outputTokens: 40, + cacheReadInputTokens: 1600, + cacheWriteInputTokens: 400 + } + } + } + ]) + ) + + expect(usage.prompt_tokens).toBe(200) + expect(usage.completion_tokens).toBe(40) + expect(usage.cache_read_input_tokens).toBe(1600) + expect(usage.cache_creation_input_tokens).toBe(400) + expect(usage.total_tokens).toBe(240) + }) + + test('emits zeros when metadataEvent.tokenUsage omits cache fields', async () => { + const usage = await readJsonUsage( + sdkStreamOf([ + { assistantResponseEvent: { content: 'Answer.' } }, + { + metadataEvent: { + tokenUsage: { inputTokens: 30, outputTokens: 5 } + } + } + ]) + ) + + expect(usage.prompt_tokens).toBe(30) + expect(usage.completion_tokens).toBe(5) + expect(usage.cache_read_input_tokens).toBe(0) + expect(usage.cache_creation_input_tokens).toBe(0) + }) + + test('emits zeros when metadataEvent has no tokenUsage', async () => { + const usage = await readJsonUsage( + sdkStreamOf([{ assistantResponseEvent: { content: 'Answer.' } }, { metadataEvent: {} }]) + ) + + expect(usage.cache_read_input_tokens).toBe(0) + expect(usage.cache_creation_input_tokens).toBe(0) + }) +}) + +describe('cache token usage — raw HTTP streaming (non-SDK)', () => { + // The non-SDK path's event shape only carries contextUsagePercentage — + // these tests pin that behavior so future contributors do not silently + // rewrite the hardcoded zeros. + function makeReadableStreamFromEvents(events: any[]) { + const encoder = new TextEncoder() + return new Response( + new ReadableStream({ + start(controller) { + for (const event of events) { + controller.enqueue(encoder.encode(JSON.stringify(event) + '\n')) + } + controller.close() + } + }) + ) + } + + async function lastUsageChunk(events: any[]): Promise { + const response = makeReadableStreamFromEvents(events) + const chunks: any[] = [] + for await (const chunk of transformKiroStream(response, MODEL, 'conversation-1')) { + chunks.push(chunk) + } + for (let i = chunks.length - 1; i >= 0; i--) { + if (chunks[i]?.usage) return chunks[i].usage + } + return null + } + + test('emits zeros — raw event stream does not carry cache token fields', async () => { + const usage = await lastUsageChunk([{ content: 'Answer.' }, { contextUsagePercentage: 5 }]) + + expect(usage).not.toBeNull() + expect(usage.cache_read_input_tokens).toBe(0) + expect(usage.cache_creation_input_tokens).toBe(0) + }) +}) diff --git a/src/core/request/response-handler.ts b/src/core/request/response-handler.ts index d1b10bf..95f19f8 100644 --- a/src/core/request/response-handler.ts +++ b/src/core/request/response-handler.ts @@ -105,10 +105,16 @@ export class ResponseHandler { finish_reason: p.stopReason === 'tool_use' ? 'tool_calls' : 'stop' } ], + // parseEventStream only surfaces `contextUsagePercentage` from the raw + // HTTP event stream — per-request cache read/write token counts are not + // available on this path. Only the SDK non-streaming path below can + // report them via MetadataEvent.tokenUsage. usage: { prompt_tokens: p.inputTokens || 0, completion_tokens: p.outputTokens || 0, - total_tokens: (p.inputTokens || 0) + (p.outputTokens || 0) + total_tokens: (p.inputTokens || 0) + (p.outputTokens || 0), + cache_creation_input_tokens: 0, + cache_read_input_tokens: 0 } } @@ -140,6 +146,8 @@ export class ResponseHandler { const toolCallOrder: string[] = [] let inputTokens = 0 let outputTokens = 0 + let cacheReadInputTokens = 0 + let cacheWriteInputTokens = 0 const eventStream = sdkResponse.generateAssistantResponseResponse if (eventStream) { @@ -169,8 +177,15 @@ export class ResponseHandler { } } if (event.metadataEvent?.tokenUsage) { - inputTokens = event.metadataEvent.tokenUsage.inputTokens || 0 - outputTokens = event.metadataEvent.tokenUsage.outputTokens || 0 + const tu = event.metadataEvent.tokenUsage + inputTokens = tu.inputTokens || 0 + outputTokens = tu.outputTokens || 0 + if (typeof tu.cacheReadInputTokens === 'number') { + cacheReadInputTokens = tu.cacheReadInputTokens + } + if (typeof tu.cacheWriteInputTokens === 'number') { + cacheWriteInputTokens = tu.cacheWriteInputTokens + } } } } @@ -197,7 +212,9 @@ export class ResponseHandler { usage: { prompt_tokens: inputTokens, completion_tokens: outputTokens, - total_tokens: inputTokens + outputTokens + total_tokens: inputTokens + outputTokens, + cache_creation_input_tokens: cacheWriteInputTokens, + cache_read_input_tokens: cacheReadInputTokens } } diff --git a/src/plugin/streaming/openai-converter.ts b/src/plugin/streaming/openai-converter.ts index 97191d2..2e555d1 100644 --- a/src/plugin/streaming/openai-converter.ts +++ b/src/plugin/streaming/openai-converter.ts @@ -64,7 +64,9 @@ export function convertToOpenAI(event: StreamEvent, id: string, model: string): ;(base as any).usage = { prompt_tokens: event.usage?.input_tokens || 0, completion_tokens: event.usage?.output_tokens || 0, - total_tokens: (event.usage?.input_tokens || 0) + (event.usage?.output_tokens || 0) + total_tokens: (event.usage?.input_tokens || 0) + (event.usage?.output_tokens || 0), + cache_creation_input_tokens: event.usage?.cache_creation_input_tokens || 0, + cache_read_input_tokens: event.usage?.cache_read_input_tokens || 0 } } else { // Skip Anthropic-specific events that @ai-sdk/openai-compatible doesn't understand diff --git a/src/plugin/streaming/sdk-stream-transformer.ts b/src/plugin/streaming/sdk-stream-transformer.ts index 877aa58..7a62dfe 100644 --- a/src/plugin/streaming/sdk-stream-transformer.ts +++ b/src/plugin/streaming/sdk-stream-transformer.ts @@ -43,6 +43,8 @@ export async function* transformSdkStream( let outputTokens = 0 let inputTokens = 0 let contextUsagePercentage: number | null = null + let cacheReadInputTokens = 0 + let cacheWriteInputTokens = 0 const toolCallFragments = new Map() const toolCallOrder: string[] = [] @@ -184,6 +186,15 @@ export async function* transformSdkStream( if (event.metadataEvent.contextUsagePercentage) { contextUsagePercentage = event.metadataEvent.contextUsagePercentage } + const tu = event.metadataEvent.tokenUsage + if (tu) { + if (typeof tu.cacheReadInputTokens === 'number') { + cacheReadInputTokens = tu.cacheReadInputTokens + } + if (typeof tu.cacheWriteInputTokens === 'number') { + cacheWriteInputTokens = tu.cacheWriteInputTokens + } + } } else if ((event as any).contextUsageEvent) { const cue = (event as any).contextUsageEvent if (cue.contextUsagePercentage) { @@ -312,8 +323,8 @@ export async function* transformSdkStream( usage: { input_tokens: inputTokens, output_tokens: outputTokens, - cache_creation_input_tokens: 0, - cache_read_input_tokens: 0 + cache_creation_input_tokens: cacheWriteInputTokens, + cache_read_input_tokens: cacheReadInputTokens } }, conversationId, diff --git a/src/plugin/streaming/stream-transformer.ts b/src/plugin/streaming/stream-transformer.ts index c2af989..48ba777 100644 --- a/src/plugin/streaming/stream-transformer.ts +++ b/src/plugin/streaming/stream-transformer.ts @@ -310,6 +310,10 @@ export async function* transformKiroStream( { type: 'message_delta', delta: { stop_reason: toolCalls.length > 0 ? 'tool_use' : 'end_turn' }, + // The non-SDK (raw HTTP) event stream carries only + // `contextUsagePercentage` — no per-request token usage is exposed, so + // cache read/write totals remain 0 here. Only the SDK path can report + // them via MetadataEvent.tokenUsage. usage: { input_tokens: inputTokens, output_tokens: outputTokens,