Skip to content

Commit c64bf5f

Browse files
committed
refactor(api): centralize service-tier primitives
1 parent d27153a commit c64bf5f

12 files changed

Lines changed: 330 additions & 59 deletions

File tree

packages/types/src/model.ts

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -54,12 +54,21 @@ export const verbosityLevelsSchema = z.enum(verbosityLevels)
5454

5555
export type VerbosityLevel = z.infer<typeof verbosityLevelsSchema>
5656

57+
/** Serialized service tier field used in provider request payloads and responses. */
58+
export const SERVICE_TIER_KEY = "service_tier"
59+
5760
/**
58-
* Service tiers (OpenAI Responses API)
61+
* Service tiers for the public OpenAI Responses API.
5962
*/
60-
export const serviceTiers = ["default", "flex", "priority"] as const
63+
export enum OpenAiServiceTier {
64+
Default = "default",
65+
Flex = "flex",
66+
Priority = "priority",
67+
}
68+
69+
export const serviceTiers = Object.values(OpenAiServiceTier) as [`${OpenAiServiceTier}`, ...`${OpenAiServiceTier}`[]]
6170
export const serviceTierSchema = z.enum(serviceTiers)
62-
export type ServiceTier = z.infer<typeof serviceTierSchema>
71+
export type ServiceTier = `${OpenAiServiceTier}`
6372

6473
/**
6574
* ModelParameter

src/api/providers/__tests__/bedrock.spec.ts

Lines changed: 9 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -58,6 +58,7 @@ import {
5858
BEDROCK_1M_CONTEXT_MODEL_IDS,
5959
BEDROCK_SERVICE_TIER_MODEL_IDS,
6060
bedrockModels,
61+
SERVICE_TIER_KEY,
6162
ApiProviderError,
6263
} from "@roo-code/types"
6364

@@ -1233,10 +1234,10 @@ describe("AwsBedrockHandler", () => {
12331234
const commandArg = mockConverseStreamCommand.mock.calls[0][0] as any
12341235

12351236
// service_tier should be at the top level of the payload
1236-
expect(commandArg.service_tier).toBe("PRIORITY")
1237+
expect(commandArg[SERVICE_TIER_KEY]).toBe("PRIORITY")
12371238
// service_tier should NOT be in additionalModelRequestFields
12381239
if (commandArg.additionalModelRequestFields) {
1239-
expect(commandArg.additionalModelRequestFields.service_tier).toBeUndefined()
1240+
expect(commandArg.additionalModelRequestFields[SERVICE_TIER_KEY]).toBeUndefined()
12401241
}
12411242
})
12421243

@@ -1263,10 +1264,10 @@ describe("AwsBedrockHandler", () => {
12631264
const commandArg = mockConverseStreamCommand.mock.calls[0][0] as any
12641265

12651266
// service_tier should be at the top level of the payload
1266-
expect(commandArg.service_tier).toBe("FLEX")
1267+
expect(commandArg[SERVICE_TIER_KEY]).toBe("FLEX")
12671268
// service_tier should NOT be in additionalModelRequestFields
12681269
if (commandArg.additionalModelRequestFields) {
1269-
expect(commandArg.additionalModelRequestFields.service_tier).toBeUndefined()
1270+
expect(commandArg.additionalModelRequestFields[SERVICE_TIER_KEY]).toBeUndefined()
12701271
}
12711272
})
12721273

@@ -1294,9 +1295,9 @@ describe("AwsBedrockHandler", () => {
12941295
const commandArg = mockConverseStreamCommand.mock.calls[0][0] as any
12951296

12961297
// Service tier should NOT be included for unsupported models (at top level or in additionalModelRequestFields)
1297-
expect(commandArg.service_tier).toBeUndefined()
1298+
expect(commandArg[SERVICE_TIER_KEY]).toBeUndefined()
12981299
if (commandArg.additionalModelRequestFields) {
1299-
expect(commandArg.additionalModelRequestFields.service_tier).toBeUndefined()
1300+
expect(commandArg.additionalModelRequestFields[SERVICE_TIER_KEY]).toBeUndefined()
13001301
}
13011302
})
13021303

@@ -1323,9 +1324,9 @@ describe("AwsBedrockHandler", () => {
13231324
const commandArg = mockConverseStreamCommand.mock.calls[0][0] as any
13241325

13251326
// Service tier should NOT be included when not specified (at top level or in additionalModelRequestFields)
1326-
expect(commandArg.service_tier).toBeUndefined()
1327+
expect(commandArg[SERVICE_TIER_KEY]).toBeUndefined()
13271328
if (commandArg.additionalModelRequestFields) {
1328-
expect(commandArg.additionalModelRequestFields.service_tier).toBeUndefined()
1329+
expect(commandArg.additionalModelRequestFields[SERVICE_TIER_KEY]).toBeUndefined()
13291330
}
13301331
})
13311332
})

src/api/providers/__tests__/openai-native-usage.spec.ts

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
import { describe, it, expect, beforeEach } from "vitest"
22
import { OpenAiNativeHandler } from "../openai-native"
3-
import { openAiNativeModels } from "@roo-code/types"
3+
import { OpenAiServiceTier, openAiNativeModels } from "@roo-code/types"
44

55
describe("OpenAiNativeHandler - normalizeUsage", () => {
66
let handler: OpenAiNativeHandler
@@ -468,7 +468,7 @@ describe("OpenAiNativeHandler - normalizeUsage", () => {
468468
it("should not apply GPT-5.4 long-context pricing to priority tier", () => {
469469
handler = new OpenAiNativeHandler({
470470
openAiNativeApiKey: "test-key",
471-
openAiNativeServiceTier: "priority",
471+
openAiNativeServiceTier: OpenAiServiceTier.Priority,
472472
})
473473

474474
const usage = {

src/api/providers/__tests__/openai-native.spec.ts

Lines changed: 68 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,7 @@ vitest.mock("@roo-code/telemetry", () => ({
1313
import { Anthropic } from "@anthropic-ai/sdk"
1414
import OpenAI from "openai"
1515

16-
import { ApiProviderError } from "@roo-code/types"
16+
import { ApiProviderError, OpenAiServiceTier, SERVICE_TIER_KEY } from "@roo-code/types"
1717

1818
import { OpenAiNativeHandler } from "../openai-native"
1919
import { ApiHandlerOptions } from "../../../shared/api"
@@ -122,6 +122,29 @@ describe("OpenAiNativeHandler", () => {
122122
})
123123

124124
describe("createMessage", () => {
125+
it.each([OpenAiServiceTier.Default, OpenAiServiceTier.Priority])(
126+
"should include the selected %s service tier",
127+
async (serviceTier) => {
128+
mockResponsesCreate.mockResolvedValue({
129+
async *[Symbol.asyncIterator]() {},
130+
})
131+
handler = new OpenAiNativeHandler({
132+
...mockOptions,
133+
apiModelId: "gpt-5.6-sol",
134+
openAiNativeServiceTier: serviceTier,
135+
})
136+
137+
for await (const chunk of handler.createMessage(systemPrompt, messages)) {
138+
void chunk
139+
}
140+
141+
expect(mockResponsesCreate).toHaveBeenCalledWith(
142+
expect.objectContaining({ [SERVICE_TIER_KEY]: serviceTier }),
143+
expect.any(Object),
144+
)
145+
},
146+
)
147+
125148
it("should handle streaming responses via Responses API", async () => {
126149
// Mock fetch for Responses API fallback
127150
const mockFetch = vitest.fn().mockResolvedValue({
@@ -221,6 +244,28 @@ describe("OpenAiNativeHandler", () => {
221244
)
222245
})
223246

247+
it.each([OpenAiServiceTier.Default, OpenAiServiceTier.Priority])(
248+
"should include the selected %s service tier",
249+
async (serviceTier) => {
250+
mockResponsesCreate.mockResolvedValue({ output: [] })
251+
handler = new OpenAiNativeHandler({
252+
...mockOptions,
253+
apiModelId: "gpt-5.6-sol",
254+
openAiNativeServiceTier: serviceTier,
255+
})
256+
257+
await handler.completePrompt("Test prompt")
258+
259+
expect(mockResponsesCreate).toHaveBeenCalledWith(
260+
expect.objectContaining({
261+
stream: false,
262+
[SERVICE_TIER_KEY]: serviceTier,
263+
}),
264+
expect.any(Object),
265+
)
266+
},
267+
)
268+
224269
it("should handle SDK errors in completePrompt", async () => {
225270
// Mock SDK to throw an error
226271
mockResponsesCreate.mockRejectedValue(new Error("API Error"))
@@ -332,12 +377,33 @@ describe("OpenAiNativeHandler", () => {
332377
expect(modelInfo.info.longContextPricing).toBeUndefined()
333378
expect(modelInfo.info.tiers).toEqual([
334379
expect.objectContaining({
335-
name: "flex",
380+
name: OpenAiServiceTier.Flex,
336381
outputPrice: 0.625,
337382
}),
338383
])
339384
})
340385

386+
it("should retain standard pricing for an explicitly selected default tier", () => {
387+
const defaultTierHandler = new OpenAiNativeHandler({
388+
...mockOptions,
389+
apiModelId: "gpt-5.4",
390+
openAiNativeServiceTier: OpenAiServiceTier.Default,
391+
})
392+
const model = defaultTierHandler.getModel()
393+
const normalizeUsage = Reflect.get(defaultTierHandler, "normalizeUsage")
394+
395+
const result = Reflect.apply(normalizeUsage, defaultTierHandler, [
396+
{
397+
input_tokens: 100_000,
398+
output_tokens: 1_000,
399+
cache_read_input_tokens: 20_000,
400+
},
401+
model,
402+
]) as { totalCost: number }
403+
404+
expect(result.totalCost).toBeCloseTo(0.22, 6)
405+
})
406+
341407
it("should return GPT-5.3 Chat model info when selected", () => {
342408
const chatHandler = new OpenAiNativeHandler({
343409
...mockOptions,

src/api/providers/bedrock.ts

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,7 @@ import {
3333
BEDROCK_GLOBAL_INFERENCE_MODEL_IDS,
3434
BEDROCK_SERVICE_TIER_MODEL_IDS,
3535
BEDROCK_SERVICE_TIER_PRICING,
36+
SERVICE_TIER_KEY,
3637
ApiProviderError,
3738
} from "@roo-code/types"
3839
import { TelemetryService } from "@roo-code/telemetry"
@@ -99,7 +100,7 @@ interface BedrockPayload {
99100
// AWS Bedrock service tiers (STANDARD, FLEX, PRIORITY) are specified at the top level
100101
// https://docs.aws.amazon.com/bedrock/latest/userguide/service-tiers-inference.html
101102
type BedrockPayloadWithServiceTier = BedrockPayload & {
102-
service_tier?: BedrockServiceTier
103+
[SERVICE_TIER_KEY]?: BedrockServiceTier
103104
}
104105

105106
// Define specific types for content block events to avoid 'as any' usage
@@ -553,7 +554,7 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH
553554
...(thinkingEnabled && { anthropic_version: "bedrock-2023-05-31" }),
554555
toolConfig,
555556
// Add service_tier as a top-level parameter (not inside additionalModelRequestFields)
556-
...(useServiceTier && { service_tier: this.options.awsBedrockServiceTier }),
557+
...(useServiceTier && { [SERVICE_TIER_KEY]: this.options.awsBedrockServiceTier }),
557558
}
558559

559560
// Create AbortController with 10 minute timeout

src/api/providers/openai-native.ts

Lines changed: 14 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,8 @@ import {
1313
type ReasoningEffort,
1414
type VerbosityLevel,
1515
type ReasoningEffortExtended,
16+
OpenAiServiceTier,
17+
SERVICE_TIER_KEY,
1618
type ServiceTier,
1719
ApiProviderError,
1820
} from "@roo-code/types"
@@ -318,7 +320,7 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio
318320
max_output_tokens?: number
319321
store?: boolean
320322
instructions?: string
321-
service_tier?: ServiceTier
323+
[SERVICE_TIER_KEY]?: ServiceTier
322324
include?: string[]
323325
/** Prompt cache retention policy: "in_memory" (default) or "24h" for extended caching */
324326
prompt_cache_retention?: "in_memory" | "24h"
@@ -369,8 +371,8 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio
369371
...(model.maxTokens ? { max_output_tokens: model.maxTokens } : {}),
370372
// Include tier when selected and supported by the model, or when explicitly "default"
371373
...(requestedTier &&
372-
(requestedTier === "default" || allowedTierNames.has(requestedTier)) && {
373-
service_tier: requestedTier,
374+
(requestedTier === OpenAiServiceTier.Default || allowedTierNames.has(requestedTier)) && {
375+
[SERVICE_TIER_KEY]: requestedTier,
374376
}),
375377
// Enable extended prompt cache retention for models that support it.
376378
// This uses the OpenAI Responses API `prompt_cache_retention` parameter.
@@ -705,8 +707,8 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio
705707
const parsed = JSON.parse(data)
706708

707709
// Capture resolved service tier if present
708-
if (parsed.response?.service_tier) {
709-
this.lastServiceTier = parsed.response.service_tier as ServiceTier
710+
if (parsed.response?.[SERVICE_TIER_KEY]) {
711+
this.lastServiceTier = parsed.response[SERVICE_TIER_KEY] as ServiceTier
710712
}
711713
// Capture complete output array (includes reasoning items with encrypted_content)
712714
if (parsed.response?.output && Array.isArray(parsed.response.output)) {
@@ -1016,8 +1018,8 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio
10161018
}
10171019
} else if (parsed.type === "response.completed" || parsed.type === "response.done") {
10181020
// Capture resolved service tier if present
1019-
if (parsed.response?.service_tier) {
1020-
this.lastServiceTier = parsed.response.service_tier as ServiceTier
1021+
if (parsed.response?.[SERVICE_TIER_KEY]) {
1022+
this.lastServiceTier = parsed.response[SERVICE_TIER_KEY] as ServiceTier
10211023
}
10221024
// Capture top-level response id
10231025
if (parsed.response?.id) {
@@ -1146,8 +1148,8 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio
11461148
*/
11471149
private async *processEvent(event: any, model: OpenAiNativeModel): ApiStream {
11481150
// Capture resolved service tier when available
1149-
if (event?.response?.service_tier) {
1150-
this.lastServiceTier = event.response.service_tier as ServiceTier
1151+
if (event?.response?.[SERVICE_TIER_KEY]) {
1152+
this.lastServiceTier = event.response[SERVICE_TIER_KEY] as ServiceTier
11511153
}
11521154
// Capture complete output array (includes reasoning items with encrypted_content)
11531155
if (event?.response?.output && Array.isArray(event.response.output)) {
@@ -1418,7 +1420,7 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio
14181420
* If no tier or no overrides exist, the original ModelInfo is returned.
14191421
*/
14201422
private applyServiceTierPricing(info: ModelInfo, tier?: ServiceTier): ModelInfo {
1421-
if (!tier || tier === "default") return info
1423+
if (!tier || tier === OpenAiServiceTier.Default) return info
14221424

14231425
// Find the tier with matching name in the tiers array
14241426
const tierInfo = info.tiers?.find((t) => t.name === tier)
@@ -1512,8 +1514,8 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio
15121514
// Include service tier if selected and supported
15131515
const requestedTier = (this.options.openAiNativeServiceTier as ServiceTier | undefined) || undefined
15141516
const allowedTierNames = new Set(model.info.tiers?.map((t) => t.name).filter(Boolean) || [])
1515-
if (requestedTier && (requestedTier === "default" || allowedTierNames.has(requestedTier))) {
1516-
requestBody.service_tier = requestedTier
1517+
if (requestedTier && (requestedTier === OpenAiServiceTier.Default || allowedTierNames.has(requestedTier))) {
1518+
requestBody[SERVICE_TIER_KEY] = requestedTier
15171519
}
15181520

15191521
// Add reasoning if supported

src/shared/cost.ts

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,4 @@
1-
import type { ModelInfo } from "@roo-code/types"
2-
import type { ServiceTier } from "@roo-code/types"
1+
import { OpenAiServiceTier, type ModelInfo, type ServiceTier } from "@roo-code/types"
32

43
export interface ApiCostResult {
54
totalInputTokens: number
@@ -13,7 +12,7 @@ function applyLongContextPricing(modelInfo: ModelInfo, totalInputTokens: number,
1312
return modelInfo
1413
}
1514

16-
const effectiveServiceTier = serviceTier ?? "default"
15+
const effectiveServiceTier = serviceTier ?? OpenAiServiceTier.Default
1716
if (pricing.appliesToServiceTiers && !pricing.appliesToServiceTiers.includes(effectiveServiceTier)) {
1817
return modelInfo
1918
}

src/utils/__tests__/cost.spec.ts

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
// npx vitest utils/__tests__/cost.spec.ts
22

3-
import type { ModelInfo } from "@roo-code/types"
3+
import { OpenAiServiceTier, type ModelInfo } from "@roo-code/types"
44

55
import { calculateApiCostAnthropic, calculateApiCostOpenAI } from "../../shared/cost"
66

@@ -283,7 +283,7 @@ describe("Cost Utility", () => {
283283
thresholdTokens: 272_000,
284284
inputPriceMultiplier: 2,
285285
outputPriceMultiplier: 1.5,
286-
appliesToServiceTiers: ["default", "flex"],
286+
appliesToServiceTiers: [OpenAiServiceTier.Default, OpenAiServiceTier.Flex],
287287
},
288288
}
289289

@@ -293,7 +293,7 @@ describe("Cost Utility", () => {
293293
1_000,
294294
undefined,
295295
100_000,
296-
"priority",
296+
OpenAiServiceTier.Priority,
297297
)
298298

299299
// Input cost: (5.0 / 1_000_000) * (300000 - 100000) = 1.0

0 commit comments

Comments
 (0)