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
88 changes: 62 additions & 26 deletions apps/web/src/lib/rewriteModelResponse.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import {
type RequestLoggingParams,
} from './rewriteModelResponse';
import { isDynamicallyOptedIntoRequestLogging } from '@/lib/ai-gateway/request-logging-opt-ins';
import { QWEN37_PLUS_MODEL_ID } from '@/lib/ai-gateway/custom-pricing';
import { KILO_ORGANIZATION_ID } from '@/lib/organizations/constants';

jest.mock('next/server', () => ({
Expand Down Expand Up @@ -102,7 +103,7 @@ describe.each(rewriters)('%s response read errors', (_name, rewrite) => {
['ResponseAborted', 'upstream_disconnect', 'disconnected'],
['TimeoutError', 'timeout', 'timed out'],
])('returns structured JSON for %s', async (errorName, errorType, messageFragment) => {
const result = await rewrite(failingResponse('application/json', errorName));
const result = await rewrite(failingResponse('application/json', errorName), true, null, null);

expect(result.status).toBe(503);
expect(await result.json()).toEqual({
Expand Down Expand Up @@ -150,7 +151,12 @@ describe.each(rewriters)('%s response read errors', (_name, rewrite) => {
});

test('omits the request id suffix when no vercel request id is available', async () => {
const result = await rewrite(failingResponse('text/event-stream', 'ResponseAborted'));
const result = await rewrite(
failingResponse('text/event-stream', 'ResponseAborted'),
true,
null,
null
);
const events = dataObjects(await readOutputStream(result)) as {
error: { message: string; vercel_request_id?: string };
}[];
Expand Down Expand Up @@ -178,7 +184,7 @@ describe('rewriteModelResponse_ChatCompletions', () => {
},
});

const result = await rewriteModelResponse_ChatCompletions(upstream);
const result = await rewriteModelResponse_ChatCompletions(upstream, true, null, null);
const json = await result.json();

expect(json.model).toBe('upstream-model');
Expand All @@ -201,7 +207,7 @@ describe('rewriteModelResponse_ChatCompletions', () => {
},
});

const result = await rewriteModelResponse_ChatCompletions(upstream);
const result = await rewriteModelResponse_ChatCompletions(upstream, true, null, null);
const json = await result.json();

expect(json.usage.prompt_tokens_details.cached_tokens).toBe(0);
Expand All @@ -214,7 +220,7 @@ describe('rewriteModelResponse_ChatCompletions', () => {
headers: { 'content-type': 'application/json' },
});

const result = await rewriteModelResponse_ChatCompletions(upstream);
const result = await rewriteModelResponse_ChatCompletions(upstream, true, null, null);

expect(result.status).toBe(502);
expect(await result.text()).toBe('not-json{');
Expand Down Expand Up @@ -242,7 +248,7 @@ describe('rewriteModelResponse_ChatCompletions', () => {
);

try {
const result = await rewriteModelResponse_ChatCompletions(upstream);
const result = await rewriteModelResponse_ChatCompletions(upstream, true, null, null);
const reader = result.body?.getReader();
expect(reader).toBeDefined();
await reader?.read();
Expand All @@ -267,7 +273,7 @@ describe('rewriteModelResponse_ChatCompletions', () => {
'data: {"id":"gen-chat","model":"upstream-model","choices":[]}\n\n'
);

const result = await rewriteModelResponse_ChatCompletions(upstream);
const result = await rewriteModelResponse_ChatCompletions(upstream, true, null, null);
const sse = await readOutputStream(result);
const events = dataObjects(sse) as Array<{ error?: { code: number; type: string } }>;

Expand All @@ -283,7 +289,7 @@ describe('rewriteModelResponse_ChatCompletions', () => {
'data: [DONE]\n\n'
);

const result = await rewriteModelResponse_ChatCompletions(upstream);
const result = await rewriteModelResponse_ChatCompletions(upstream, true, null, null);
const sse = await readOutputStream(result);
const [chunk] = dataObjects(sse) as Array<{
model: string;
Expand All @@ -301,7 +307,7 @@ describe('rewriteModelResponse_ChatCompletions', () => {
'data: {"model":"upstream-model","usage":{"cost":1,"is_byok":true,"prompt_tokens":4,"completion_tokens":2,"total_tokens":6,"prompt_tokens_details":{}}}\n\n'
);

const result = await rewriteModelResponse_ChatCompletions(upstream);
const result = await rewriteModelResponse_ChatCompletions(upstream, true, null, null);
const sse = await readOutputStream(result);
const [chunk] = dataObjects(sse) as Array<{
model: string;
Expand All @@ -325,7 +331,7 @@ describe('rewriteModelResponse_ChatCompletions', () => {
': openrouter heartbeat\n\n' + 'data: {"model":"upstream-model","choices":[]}\n\n'
);

const result = await rewriteModelResponse_ChatCompletions(upstream);
const result = await rewriteModelResponse_ChatCompletions(upstream, true, null, null);
const sse = await readOutputStream(result);

expect(sse).toContain(': KILO PROCESSING');
Expand All @@ -337,7 +343,7 @@ describe('rewriteModelResponse_ChatCompletions', () => {
headers: { 'content-type': 'text/event-stream' },
});

const result = await rewriteModelResponse_ChatCompletions(upstream);
const result = await rewriteModelResponse_ChatCompletions(upstream, true, null, null);

expect(await readOutputStream(result)).toBe('');
});
Expand All @@ -354,7 +360,10 @@ describe('rewriteModelResponse_Messages', () => {
'text/event-stream',
errorName,
'data: {"type":"message_start","message":{"id":"gen-message","usage":{"input_tokens":1,"output_tokens":0}}}\n\n'
)
),
true,
null,
null
);
const sse = await readOutputStream(result);

Expand Down Expand Up @@ -392,7 +401,7 @@ describe('rewriteModelResponse_Messages', () => {
},
});

const result = await rewriteModelResponse_Messages(upstream);
const result = await rewriteModelResponse_Messages(upstream, true, null, null);
const json = await result.json();

expect(json.model).toBe('upstream-model');
Expand All @@ -408,7 +417,7 @@ describe('rewriteModelResponse_Messages', () => {
headers: { 'content-type': 'application/json' },
});

const result = await rewriteModelResponse_Messages(upstream);
const result = await rewriteModelResponse_Messages(upstream, true, null, null);

expect(result.status).toBe(500);
expect(await result.text()).toBe('}{');
Expand All @@ -421,7 +430,7 @@ describe('rewriteModelResponse_Messages', () => {
'data: [DONE]\n\n'
);

const result = await rewriteModelResponse_Messages(upstream);
const result = await rewriteModelResponse_Messages(upstream, true, null, null);
const sse = await readOutputStream(result);
const events = dataObjects(sse) as Array<{
type: string;
Expand Down Expand Up @@ -450,7 +459,7 @@ describe('rewriteModelResponse_Messages', () => {
'data: {"type":"message_delta","usage":{"output_tokens":9},"delta":{}}\n\n'
);

const result = await rewriteModelResponse_Messages(upstream);
const result = await rewriteModelResponse_Messages(upstream, true, null, null);
const sse = await readOutputStream(result);

expect(dataPayloads(sse)).not.toContain('[DONE]');
Expand All @@ -467,7 +476,10 @@ describe('rewriteModelResponse_Responses', () => {
'text/event-stream',
errorName,
'data: {"type":"response.created","sequence_number":4,"response":{"id":"gen-response"}}\n\n'
)
),
true,
null,
null
);
const sse = await readOutputStream(result);

Expand Down Expand Up @@ -505,7 +517,7 @@ describe('rewriteModelResponse_Responses', () => {
},
});

const result = await rewriteModelResponse_Responses(upstream);
const result = await rewriteModelResponse_Responses(upstream, true, null, null);
const json = await result.json();

expect(json.model).toBe('upstream-model');
Expand All @@ -521,7 +533,7 @@ describe('rewriteModelResponse_Responses', () => {
'data: [DONE]\n\n'
);

const result = await rewriteModelResponse_Responses(upstream);
const result = await rewriteModelResponse_Responses(upstream, true, null, null);
const sse = await readOutputStream(result);
const [event] = dataObjects(sse) as Array<{
type: string;
Expand Down Expand Up @@ -626,6 +638,26 @@ describe('rewriteModelResponse', () => {
});
});

test('strips cost for models with custom pricing', async () => {
const result = await rewriteModelResponse(
jsonResponse({
model: QWEN37_PLUS_MODEL_ID,
usage: { cost: 0.5, cost_details: { upstream_inference_cost: 0.4 }, is_byok: false },
}),
QWEN37_PLUS_MODEL_ID,
'openrouter',
'chat_completions',
makeLogging()
);

// The upstream-reported cost does not reflect the custom pricing, so it
// must be removed just like for free models.
expect(await result.json()).toEqual({
model: QWEN37_PLUS_MODEL_ID,
usage: {},
});
});

test('processes paid-model responses when request logging is enabled', async () => {
mockedOptIn.mockResolvedValueOnce(true);
const result = await rewriteModelResponse(
Expand All @@ -649,7 +681,7 @@ describe('request log capture', () => {
const capture = makeCapture();
const body = { model: 'upstream-model' };

const result = await rewrite(jsonResponse(body), true, capture);
const result = await rewrite(jsonResponse(body), true, capture, null);

expect(result.status).toBe(200);
expect(capture.setBody).toHaveBeenCalledTimes(1);
Expand All @@ -662,7 +694,7 @@ describe('request log capture', () => {
const sseBody =
'data: {"id":"gen-1","model":"upstream-model","choices":[]}\n\n' + 'data: [DONE]\n\n';

const result = await rewrite(sseResponse(sseBody), true, capture);
const result = await rewrite(sseResponse(sseBody), true, capture, null);
await readOutputStream(result);

expect(capture.setBody).toHaveBeenCalledTimes(1);
Expand All @@ -678,7 +710,8 @@ describe('request log capture', () => {
const result = await rewrite(
new Response(null, { headers: { 'content-type': 'text/event-stream' } }),
true,
capture
capture,
null
);
await readOutputStream(result);

Expand All @@ -696,7 +729,8 @@ describe('request log capture', () => {
const result = await rewrite(
failingResponse('text/event-stream', 'ResponseAborted', receivedChunks),
true,
capture
capture,
null
);
await readOutputStream(result);

Expand All @@ -714,7 +748,8 @@ describe('request log capture', () => {
const result = await rewrite(
failingResponse('text/event-stream', 'ResponseAborted'),
true,
capture
capture,
null
);
await readOutputStream(result);

Expand All @@ -732,7 +767,8 @@ describe('request log capture', () => {
const result = await rewrite(
failingResponse('application/json', 'TimeoutError'),
true,
capture
capture,
null
);

expect(result.status).toBe(503);
Expand All @@ -747,7 +783,7 @@ describe('request log capture', () => {
headers: { 'content-type': 'text/event-stream' },
});

const result = await rewriteModelResponse_ChatCompletions(upstream, true, capture);
const result = await rewriteModelResponse_ChatCompletions(upstream, true, capture, null);
const reader = result.body?.getReader();
await reader?.cancel();

Expand Down
44 changes: 18 additions & 26 deletions apps/web/src/lib/rewriteModelResponse.ts
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import { api_request_log, type User } from '@kilocode/db/schema';
import { isKiloExclusiveFreeModel } from '@/lib/ai-gateway/models';
import { getCustomPricing } from '@/lib/ai-gateway/custom-pricing';
import { detectToolCallArgumentErrors } from '@/lib/ai-gateway/api-request-log-errors';
import type { GatewayRequest } from '@/lib/ai-gateway/providers/openrouter/types';
import type { ProviderId } from '@/lib/ai-gateway/providers/types';
Expand Down Expand Up @@ -218,7 +219,7 @@ async function readResponseText(
response: Response,
headers: Headers,
vercelRequestId: string | null | undefined,
capture?: RequestLogCapture | null
capture: RequestLogCapture | null
): Promise<{ text: string } | { error: unknown; errorResponse: NextResponse }> {
try {
return { text: await response.text() };
Expand Down Expand Up @@ -256,7 +257,7 @@ async function rewriteSseStream(
serializeError: (error: ResponseReadError) => string,
onFinally: () => void,
vercelRequestId: string | null | undefined,
capture?: RequestLogCapture | null
capture: RequestLogCapture | null
) {
const decoder = new TextDecoder();
// Accumulate the raw upstream text for request logging while the stream is
Expand Down Expand Up @@ -323,9 +324,9 @@ function rewriteUsage(usage: OpenRouterUsage, removeCost: boolean) {

export async function rewriteModelResponse_ChatCompletions(
response: Response,
removeCost = true,
capture?: RequestLogCapture | null,
vercelRequestId?: string | null
removeCost: boolean,
capture: RequestLogCapture | null,
vercelRequestId: string | null
) {
const headers = getOutputHeaders(response);

Expand Down Expand Up @@ -478,9 +479,9 @@ function rewriteMessagesUsage(usage: MessagesApiUsage, removeCost: boolean) {

export async function rewriteModelResponse_Messages(
response: Response,
removeCost = true,
capture?: RequestLogCapture | null,
vercelRequestId?: string | null
removeCost: boolean,
capture: RequestLogCapture | null,
vercelRequestId: string | null
) {
const headers = getOutputHeaders(response);

Expand Down Expand Up @@ -614,9 +615,9 @@ type ResponsesApiEvent = {

export async function rewriteModelResponse_Responses(
response: Response,
removeCost = true,
capture?: RequestLogCapture | null,
vercelRequestId?: string | null
removeCost: boolean,
capture: RequestLogCapture | null,
vercelRequestId: string | null
) {
const headers = getOutputHeaders(response);

Expand Down Expand Up @@ -742,34 +743,25 @@ export async function rewriteModelResponse(
logging: RequestLoggingParams
): Promise<NextResponse> {
const capture = await createRequestLogCapture(response, model, providerId, logging);
const isFreeModelRequiringCostRemoval =
(providerId === 'openrouter' || providerId === 'vercel') && isKiloExclusiveFreeModel(model);
const requiresCostRemoval =
(providerId === 'openrouter' || providerId === 'vercel') &&
(isKiloExclusiveFreeModel(model) || getCustomPricing(model) !== undefined);

console.debug('[rewriteModelResponse] rewriting response for %s', model);
const { vercel_request_id: vercelRequestId } = logging;
if (kind === 'chat_completions') {
return rewriteModelResponse_ChatCompletions(
response,
isFreeModelRequiringCostRemoval,
requiresCostRemoval,
capture,
vercelRequestId
);
}
if (kind === 'responses') {
return rewriteModelResponse_Responses(
response,
isFreeModelRequiringCostRemoval,
capture,
vercelRequestId
);
return rewriteModelResponse_Responses(response, requiresCostRemoval, capture, vercelRequestId);
}
if (kind === 'messages') {
return rewriteModelResponse_Messages(
response,
isFreeModelRequiringCostRemoval,
capture,
vercelRequestId
);
return rewriteModelResponse_Messages(response, requiresCostRemoval, capture, vercelRequestId);
}

const error = new Error(`implementation error: unrecognized API kind ${kind}`);
Expand Down