diff --git a/packages/types/package.json b/packages/types/package.json index c4cbd4454c..98cf9f373a 100644 --- a/packages/types/package.json +++ b/packages/types/package.json @@ -11,6 +11,14 @@ "types": "./dist/index.d.cts", "default": "./dist/index.cjs" } + }, + "./model": { + "types": "./src/model.ts", + "import": "./src/model.ts" + }, + "./provider-identifiers": { + "types": "./src/provider-identifiers.ts", + "import": "./src/provider-identifiers.ts" } }, "scripts": { diff --git a/packages/types/src/model.ts b/packages/types/src/model.ts index 257a30f3e7..28a32eecb4 100644 --- a/packages/types/src/model.ts +++ b/packages/types/src/model.ts @@ -54,10 +54,19 @@ export const verbosityLevelsSchema = z.enum(verbosityLevels) export type VerbosityLevel = z.infer +/** Serialized service tier field used in provider request payloads and responses. */ +export const SERVICE_TIER_KEY = "service_tier" + /** - * Service tiers (OpenAI Responses API) + * Service tiers for the public OpenAI Responses API. */ -export const serviceTiers = ["default", "flex", "priority"] as const +export const OpenAiServiceTier = { + Default: "default", + Flex: "flex", + Priority: "priority", +} as const + +export const serviceTiers = [OpenAiServiceTier.Default, OpenAiServiceTier.Flex, OpenAiServiceTier.Priority] as const export const serviceTierSchema = z.enum(serviceTiers) export type ServiceTier = z.infer diff --git a/src/api/providers/__tests__/bedrock.spec.ts b/src/api/providers/__tests__/bedrock.spec.ts index 71c2d1cd0c..b025f33f02 100644 --- a/src/api/providers/__tests__/bedrock.spec.ts +++ b/src/api/providers/__tests__/bedrock.spec.ts @@ -58,6 +58,7 @@ import { BEDROCK_1M_CONTEXT_MODEL_IDS, BEDROCK_SERVICE_TIER_MODEL_IDS, bedrockModels, + SERVICE_TIER_KEY, ApiProviderError, } from "@roo-code/types" @@ -1233,10 +1234,10 @@ describe("AwsBedrockHandler", () => { const commandArg = mockConverseStreamCommand.mock.calls[0][0] as any // service_tier should be at the top level of the payload - expect(commandArg.service_tier).toBe("PRIORITY") + expect(commandArg[SERVICE_TIER_KEY]).toBe("PRIORITY") // service_tier should NOT be in additionalModelRequestFields if (commandArg.additionalModelRequestFields) { - expect(commandArg.additionalModelRequestFields.service_tier).toBeUndefined() + expect(commandArg.additionalModelRequestFields[SERVICE_TIER_KEY]).toBeUndefined() } }) @@ -1263,10 +1264,10 @@ describe("AwsBedrockHandler", () => { const commandArg = mockConverseStreamCommand.mock.calls[0][0] as any // service_tier should be at the top level of the payload - expect(commandArg.service_tier).toBe("FLEX") + expect(commandArg[SERVICE_TIER_KEY]).toBe("FLEX") // service_tier should NOT be in additionalModelRequestFields if (commandArg.additionalModelRequestFields) { - expect(commandArg.additionalModelRequestFields.service_tier).toBeUndefined() + expect(commandArg.additionalModelRequestFields[SERVICE_TIER_KEY]).toBeUndefined() } }) @@ -1294,9 +1295,9 @@ describe("AwsBedrockHandler", () => { const commandArg = mockConverseStreamCommand.mock.calls[0][0] as any // Service tier should NOT be included for unsupported models (at top level or in additionalModelRequestFields) - expect(commandArg.service_tier).toBeUndefined() + expect(commandArg[SERVICE_TIER_KEY]).toBeUndefined() if (commandArg.additionalModelRequestFields) { - expect(commandArg.additionalModelRequestFields.service_tier).toBeUndefined() + expect(commandArg.additionalModelRequestFields[SERVICE_TIER_KEY]).toBeUndefined() } }) @@ -1323,9 +1324,9 @@ describe("AwsBedrockHandler", () => { const commandArg = mockConverseStreamCommand.mock.calls[0][0] as any // Service tier should NOT be included when not specified (at top level or in additionalModelRequestFields) - expect(commandArg.service_tier).toBeUndefined() + expect(commandArg[SERVICE_TIER_KEY]).toBeUndefined() if (commandArg.additionalModelRequestFields) { - expect(commandArg.additionalModelRequestFields.service_tier).toBeUndefined() + expect(commandArg.additionalModelRequestFields[SERVICE_TIER_KEY]).toBeUndefined() } }) }) diff --git a/src/api/providers/__tests__/openai-native-usage.spec.ts b/src/api/providers/__tests__/openai-native-usage.spec.ts index a266642e7a..184f04ec80 100644 --- a/src/api/providers/__tests__/openai-native-usage.spec.ts +++ b/src/api/providers/__tests__/openai-native-usage.spec.ts @@ -1,6 +1,6 @@ import { describe, it, expect, beforeEach } from "vitest" import { OpenAiNativeHandler } from "../openai-native" -import { openAiNativeModels } from "@roo-code/types" +import { OpenAiServiceTier, openAiNativeModels } from "@roo-code/types" describe("OpenAiNativeHandler - normalizeUsage", () => { let handler: OpenAiNativeHandler @@ -468,7 +468,7 @@ describe("OpenAiNativeHandler - normalizeUsage", () => { it("should not apply GPT-5.4 long-context pricing to priority tier", () => { handler = new OpenAiNativeHandler({ openAiNativeApiKey: "test-key", - openAiNativeServiceTier: "priority", + openAiNativeServiceTier: OpenAiServiceTier.Priority, }) const usage = { diff --git a/src/api/providers/__tests__/openai-native.spec.ts b/src/api/providers/__tests__/openai-native.spec.ts index 1acb4101be..cba0779191 100644 --- a/src/api/providers/__tests__/openai-native.spec.ts +++ b/src/api/providers/__tests__/openai-native.spec.ts @@ -13,7 +13,7 @@ vitest.mock("@roo-code/telemetry", () => ({ import { Anthropic } from "@anthropic-ai/sdk" import OpenAI from "openai" -import { ApiProviderError } from "@roo-code/types" +import { ApiProviderError, OpenAiServiceTier, SERVICE_TIER_KEY, serviceTiers } from "@roo-code/types" import { OpenAiNativeHandler } from "../openai-native" import { ApiHandlerOptions } from "../../../shared/api" @@ -22,6 +22,24 @@ import { Package } from "../../../shared/package" // Mock OpenAI client - now everything uses Responses API const mockResponsesCreate = vitest.fn() +const serviceTierPricingCases = [ + { + requestedTier: OpenAiServiceTier.Default, + resolvedTier: OpenAiServiceTier.Priority, + expectedCost: 0.00275, + }, + { + requestedTier: OpenAiServiceTier.Priority, + resolvedTier: OpenAiServiceTier.Flex, + expectedCost: 0.00055, + }, + { + requestedTier: OpenAiServiceTier.Flex, + resolvedTier: OpenAiServiceTier.Default, + expectedCost: 0.0011, + }, +] + vitest.mock("openai", () => { return { __esModule: true, @@ -122,6 +140,210 @@ describe("OpenAiNativeHandler", () => { }) describe("createMessage", () => { + it.each(serviceTiers)("should include the selected %s service tier", async (serviceTier) => { + mockResponsesCreate.mockResolvedValue({ + async *[Symbol.asyncIterator]() {}, + }) + handler = new OpenAiNativeHandler({ + ...mockOptions, + apiModelId: "gpt-5.6-sol", + openAiNativeServiceTier: serviceTier, + }) + + for await (const chunk of handler.createMessage(systemPrompt, messages)) { + void chunk + } + + expect(mockResponsesCreate).toHaveBeenCalledWith( + expect.objectContaining({ [SERVICE_TIER_KEY]: serviceTier }), + expect.any(Object), + ) + }) + + it.each(serviceTierPricingCases)( + "prices SDK stream usage using resolved $resolvedTier tier instead of requested $requestedTier tier", + async ({ requestedTier, resolvedTier, expectedCost }) => { + mockResponsesCreate.mockResolvedValue({ + async *[Symbol.asyncIterator]() { + yield { + type: "response.done", + response: { + [SERVICE_TIER_KEY]: resolvedTier, + usage: { input_tokens: 100, output_tokens: 20 }, + }, + } + }, + }) + handler = new OpenAiNativeHandler({ + ...mockOptions, + apiModelId: "gpt-5.6-sol", + openAiNativeServiceTier: requestedTier, + }) + + const chunks = [] + for await (const chunk of handler.createMessage(systemPrompt, messages)) { + chunks.push(chunk) + } + + expect(chunks).toContainEqual( + expect.objectContaining({ + type: "usage", + inputTokens: 100, + outputTokens: 20, + totalCost: expectedCost, + }), + ) + }, + ) + + it.each([ + { + name: "an explicitly selected default tier", + modelId: "gpt-5.4" as const, + requestedTier: OpenAiServiceTier.Default, + resolvedTier: undefined, + expectedCost: 0.22, + }, + { + name: "no selected service tier", + modelId: "gpt-5.4" as const, + requestedTier: undefined, + resolvedTier: undefined, + expectedCost: 0.22, + }, + { + name: "a resolved service tier without a pricing entry", + modelId: "gpt-5.6-luna" as const, + requestedTier: OpenAiServiceTier.Default, + resolvedTier: OpenAiServiceTier.Priority, + expectedCost: 0.088, + }, + ])("retains standard pricing for $name", async ({ modelId, requestedTier, resolvedTier, expectedCost }) => { + mockResponsesCreate.mockResolvedValue({ + async *[Symbol.asyncIterator]() { + yield { + type: "response.done", + response: { + ...(resolvedTier ? { [SERVICE_TIER_KEY]: resolvedTier } : {}), + usage: { + input_tokens: 100_000, + output_tokens: 1_000, + cache_read_input_tokens: 20_000, + }, + }, + } + }, + }) + handler = new OpenAiNativeHandler({ + ...mockOptions, + apiModelId: modelId, + openAiNativeServiceTier: requestedTier, + }) + + const chunks = [] + for await (const chunk of handler.createMessage(systemPrompt, messages)) { + chunks.push(chunk) + } + + const usageChunk = chunks.find((chunk) => chunk.type === "usage") + expect(usageChunk).toBeDefined() + expect(usageChunk?.totalCost).toBeCloseTo(expectedCost, 6) + }) + + it.each(serviceTierPricingCases)( + "requests $requestedTier but prices manual SSE fallback usage using OpenAI's resolved $resolvedTier tier", + async ({ requestedTier, resolvedTier, expectedCost }) => { + mockResponsesCreate.mockRejectedValue(new Error("SDK not available")) + const mockFetch = vitest.fn().mockResolvedValue({ + ok: true, + body: new ReadableStream({ + start(controller) { + controller.enqueue( + new TextEncoder().encode( + `data: ${JSON.stringify({ + type: "response.done", + response: { + [SERVICE_TIER_KEY]: resolvedTier, + usage: { input_tokens: 100, output_tokens: 20 }, + }, + })}\n\n`, + ), + ) + controller.enqueue(new TextEncoder().encode("data: [DONE]\n\n")) + controller.close() + }, + }), + }) + global.fetch = mockFetch as typeof fetch + handler = new OpenAiNativeHandler({ + ...mockOptions, + apiModelId: "gpt-5.6-sol", + openAiNativeServiceTier: requestedTier, + }) + + const chunks = [] + for await (const chunk of handler.createMessage(systemPrompt, messages)) { + chunks.push(chunk) + } + + const [, request] = mockFetch.mock.calls[0] + expect(JSON.parse(request.body)).toMatchObject({ [SERVICE_TIER_KEY]: requestedTier }) + expect(chunks).toContainEqual(expect.objectContaining({ type: "usage", totalCost: expectedCost })) + }, + ) + + it.each(serviceTierPricingCases)( + "captures resolved $resolvedTier tier from a manual SSE completion event when $requestedTier was requested", + async ({ requestedTier, resolvedTier, expectedCost }) => { + mockResponsesCreate.mockRejectedValue(new Error("SDK not available")) + const mockFetch = vitest.fn().mockResolvedValue({ + ok: true, + body: new ReadableStream({ + start(controller) { + controller.enqueue( + new TextEncoder().encode( + `data: ${JSON.stringify({ + type: "response.completed", + response: { [SERVICE_TIER_KEY]: resolvedTier }, + })}\n\n`, + ), + ) + controller.enqueue( + new TextEncoder().encode( + `data: ${JSON.stringify({ + type: "response.usage", + usage: { input_tokens: 100, output_tokens: 20 }, + })}\n\n`, + ), + ) + controller.enqueue(new TextEncoder().encode("data: [DONE]\n\n")) + controller.close() + }, + }), + }) + global.fetch = mockFetch as typeof fetch + handler = new OpenAiNativeHandler({ + ...mockOptions, + apiModelId: "gpt-5.6-sol", + openAiNativeServiceTier: requestedTier, + }) + + const chunks = [] + for await (const chunk of handler.createMessage(systemPrompt, messages)) { + chunks.push(chunk) + } + + expect(chunks).toContainEqual( + expect.objectContaining({ + type: "usage", + inputTokens: 100, + outputTokens: 20, + totalCost: expectedCost, + }), + ) + }, + ) + it("should handle streaming responses via Responses API", async () => { // Mock fetch for Responses API fallback const mockFetch = vitest.fn().mockResolvedValue({ @@ -221,6 +443,50 @@ describe("OpenAiNativeHandler", () => { ) }) + it.each(serviceTiers)("should include the selected %s service tier", async (serviceTier) => { + mockResponsesCreate.mockResolvedValue({ output: [] }) + handler = new OpenAiNativeHandler({ + ...mockOptions, + apiModelId: "gpt-5.6-sol", + openAiNativeServiceTier: serviceTier, + }) + + await handler.completePrompt("Test prompt") + + expect(mockResponsesCreate).toHaveBeenCalledWith( + expect.objectContaining({ + stream: false, + [SERVICE_TIER_KEY]: serviceTier, + }), + expect.any(Object), + ) + }) + + it("should omit the service tier when none is configured", async () => { + mockResponsesCreate.mockResolvedValue({ output: [] }) + + await handler.completePrompt("Test prompt") + + const [request] = mockResponsesCreate.mock.calls[0] + expect(request.stream).toBe(false) + expect(request).not.toHaveProperty(SERVICE_TIER_KEY) + }) + + it("should omit a configured service tier that the model does not support", async () => { + mockResponsesCreate.mockResolvedValue({ output: [] }) + handler = new OpenAiNativeHandler({ + ...mockOptions, + apiModelId: "gpt-5.6-luna", + openAiNativeServiceTier: OpenAiServiceTier.Priority, + }) + + await handler.completePrompt("Test prompt") + + const [request] = mockResponsesCreate.mock.calls[0] + expect(request.stream).toBe(false) + expect(request).not.toHaveProperty(SERVICE_TIER_KEY) + }) + it("should handle SDK errors in completePrompt", async () => { // Mock SDK to throw an error mockResponsesCreate.mockRejectedValue(new Error("API Error")) @@ -332,7 +598,7 @@ describe("OpenAiNativeHandler", () => { expect(modelInfo.info.longContextPricing).toBeUndefined() expect(modelInfo.info.tiers).toEqual([ expect.objectContaining({ - name: "flex", + name: OpenAiServiceTier.Flex, outputPrice: 0.625, }), ]) diff --git a/src/api/providers/bedrock.ts b/src/api/providers/bedrock.ts index 645ad8d354..0d39e843c4 100644 --- a/src/api/providers/bedrock.ts +++ b/src/api/providers/bedrock.ts @@ -33,6 +33,7 @@ import { BEDROCK_GLOBAL_INFERENCE_MODEL_IDS, BEDROCK_SERVICE_TIER_MODEL_IDS, BEDROCK_SERVICE_TIER_PRICING, + SERVICE_TIER_KEY, ApiProviderError, } from "@roo-code/types" import { TelemetryService } from "@roo-code/telemetry" @@ -99,7 +100,7 @@ interface BedrockPayload { // AWS Bedrock service tiers (STANDARD, FLEX, PRIORITY) are specified at the top level // https://docs.aws.amazon.com/bedrock/latest/userguide/service-tiers-inference.html type BedrockPayloadWithServiceTier = BedrockPayload & { - service_tier?: BedrockServiceTier + [SERVICE_TIER_KEY]?: BedrockServiceTier } // Define specific types for content block events to avoid 'as any' usage @@ -553,7 +554,7 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH ...(thinkingEnabled && { anthropic_version: "bedrock-2023-05-31" }), toolConfig, // Add service_tier as a top-level parameter (not inside additionalModelRequestFields) - ...(useServiceTier && { service_tier: this.options.awsBedrockServiceTier }), + ...(useServiceTier && { [SERVICE_TIER_KEY]: this.options.awsBedrockServiceTier }), } // Create AbortController with 10 minute timeout diff --git a/src/api/providers/openai-native.ts b/src/api/providers/openai-native.ts index a1ce2d89d0..8dffe03dcc 100644 --- a/src/api/providers/openai-native.ts +++ b/src/api/providers/openai-native.ts @@ -13,6 +13,8 @@ import { type ReasoningEffort, type VerbosityLevel, type ReasoningEffortExtended, + OpenAiServiceTier, + SERVICE_TIER_KEY, type ServiceTier, ApiProviderError, } from "@roo-code/types" @@ -318,7 +320,7 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio max_output_tokens?: number store?: boolean instructions?: string - service_tier?: ServiceTier + [SERVICE_TIER_KEY]?: ServiceTier include?: string[] /** Prompt cache retention policy: "in_memory" (default) or "24h" for extended caching */ prompt_cache_retention?: "in_memory" | "24h" @@ -334,8 +336,7 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio } // Validate requested tier against model support; if not supported, omit. - const requestedTier = (this.options.openAiNativeServiceTier as ServiceTier | undefined) || undefined - const allowedTierNames = new Set(model.info.tiers?.map((t) => t.name).filter(Boolean) || []) + const serviceTier = this.getAllowedServiceTier(model) // Decide whether to enable extended prompt cache retention for this request const promptCacheRetention = this.getPromptCacheRetention(model) @@ -368,10 +369,7 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio // Use the per-request reserved output computed by Roo (params.maxTokens from getModelParams). ...(model.maxTokens ? { max_output_tokens: model.maxTokens } : {}), // Include tier when selected and supported by the model, or when explicitly "default" - ...(requestedTier && - (requestedTier === "default" || allowedTierNames.has(requestedTier)) && { - service_tier: requestedTier, - }), + ...(serviceTier && { [SERVICE_TIER_KEY]: serviceTier }), // Enable extended prompt cache retention for models that support it. // This uses the OpenAI Responses API `prompt_cache_retention` parameter. ...(promptCacheRetention ? { prompt_cache_retention: promptCacheRetention } : {}), @@ -705,8 +703,8 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio const parsed = JSON.parse(data) // Capture resolved service tier if present - if (parsed.response?.service_tier) { - this.lastServiceTier = parsed.response.service_tier as ServiceTier + if (parsed.response?.[SERVICE_TIER_KEY]) { + this.lastServiceTier = parsed.response[SERVICE_TIER_KEY] as ServiceTier } // Capture complete output array (includes reasoning items with encrypted_content) if (parsed.response?.output && Array.isArray(parsed.response.output)) { @@ -1015,10 +1013,6 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio ) } } else if (parsed.type === "response.completed" || parsed.type === "response.done") { - // Capture resolved service tier if present - if (parsed.response?.service_tier) { - this.lastServiceTier = parsed.response.service_tier as ServiceTier - } // Capture top-level response id if (parsed.response?.id) { this.lastResponseId = parsed.response.id as string @@ -1146,8 +1140,8 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio */ private async *processEvent(event: any, model: OpenAiNativeModel): ApiStream { // Capture resolved service tier when available - if (event?.response?.service_tier) { - this.lastServiceTier = event.response.service_tier as ServiceTier + if (event?.response?.[SERVICE_TIER_KEY]) { + this.lastServiceTier = event.response[SERVICE_TIER_KEY] as ServiceTier } // Capture complete output array (includes reasoning items with encrypted_content) if (event?.response?.output && Array.isArray(event.response.output)) { @@ -1418,7 +1412,7 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio * If no tier or no overrides exist, the original ModelInfo is returned. */ private applyServiceTierPricing(info: ModelInfo, tier?: ServiceTier): ModelInfo { - if (!tier || tier === "default") return info + if (!tier || tier === OpenAiServiceTier.Default) return info // Find the tier with matching name in the tiers array const tierInfo = info.tiers?.find((t) => t.name === tier) @@ -1433,6 +1427,15 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio } } + private getAllowedServiceTier(model: OpenAiNativeModel): ServiceTier | undefined { + const requestedTier = this.options.openAiNativeServiceTier + const allowedTierNames = new Set(model.info.tiers?.map(({ name }) => name).filter(Boolean)) + + return requestedTier === OpenAiServiceTier.Default || (requestedTier && allowedTierNames.has(requestedTier)) + ? requestedTier + : undefined + } + // Removed isResponsesApiModel method as ALL models now use the Responses API override getModel() { @@ -1510,10 +1513,9 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio } // Include service tier if selected and supported - const requestedTier = (this.options.openAiNativeServiceTier as ServiceTier | undefined) || undefined - const allowedTierNames = new Set(model.info.tiers?.map((t) => t.name).filter(Boolean) || []) - if (requestedTier && (requestedTier === "default" || allowedTierNames.has(requestedTier))) { - requestBody.service_tier = requestedTier + const serviceTier = this.getAllowedServiceTier(model) + if (serviceTier) { + requestBody[SERVICE_TIER_KEY] = serviceTier } // Add reasoning if supported diff --git a/src/shared/cost.ts b/src/shared/cost.ts index 8954904fda..754a2d84df 100644 --- a/src/shared/cost.ts +++ b/src/shared/cost.ts @@ -1,5 +1,4 @@ -import type { ModelInfo } from "@roo-code/types" -import type { ServiceTier } from "@roo-code/types" +import { OpenAiServiceTier, type ModelInfo, type ServiceTier } from "@roo-code/types" export interface ApiCostResult { totalInputTokens: number @@ -13,7 +12,7 @@ function applyLongContextPricing(modelInfo: ModelInfo, totalInputTokens: number, return modelInfo } - const effectiveServiceTier = serviceTier ?? "default" + const effectiveServiceTier = serviceTier ?? OpenAiServiceTier.Default if (pricing.appliesToServiceTiers && !pricing.appliesToServiceTiers.includes(effectiveServiceTier)) { return modelInfo } diff --git a/src/utils/__tests__/cost.spec.ts b/src/utils/__tests__/cost.spec.ts index 6f0b594c8d..b369a0405b 100644 --- a/src/utils/__tests__/cost.spec.ts +++ b/src/utils/__tests__/cost.spec.ts @@ -1,6 +1,6 @@ // npx vitest utils/__tests__/cost.spec.ts -import type { ModelInfo } from "@roo-code/types" +import { OpenAiServiceTier, type ModelInfo } from "@roo-code/types" import { calculateApiCostAnthropic, calculateApiCostOpenAI } from "../../shared/cost" @@ -283,7 +283,7 @@ describe("Cost Utility", () => { thresholdTokens: 272_000, inputPriceMultiplier: 2, outputPriceMultiplier: 1.5, - appliesToServiceTiers: ["default", "flex"], + appliesToServiceTiers: [OpenAiServiceTier.Default, OpenAiServiceTier.Flex], }, } @@ -293,7 +293,7 @@ describe("Cost Utility", () => { 1_000, undefined, 100_000, - "priority", + OpenAiServiceTier.Priority, ) // Input cost: (5.0 / 1_000_000) * (300000 - 100000) = 1.0 diff --git a/webview-ui/playwright-ct.config.ts b/webview-ui/playwright-ct.config.ts index 8b236781b2..3eb0abac7b 100644 --- a/webview-ui/playwright-ct.config.ts +++ b/webview-ui/playwright-ct.config.ts @@ -58,10 +58,15 @@ export default defineConfig({ ], resolve: { alias: { + "@src/i18n/TranslationContext": path.resolve(dirname, "./playwright/TranslationContext.ts"), "@": path.resolve(dirname, "./src"), "@src": path.resolve(dirname, "./src"), "@roo": path.resolve(dirname, "../src/shared"), - vscode: path.resolve(dirname, "./src/__mocks__/vscode.ts"), + "@vscode/webview-ui-toolkit/react": path.resolve( + dirname, + "./src/__mocks__/@vscode/webview-ui-toolkit/react.tsx", + ), + vscode: path.resolve(dirname, "../src/__mocks__/vscode.js"), }, }, define: { diff --git a/webview-ui/playwright/TranslationContext.ts b/webview-ui/playwright/TranslationContext.ts new file mode 100644 index 0000000000..b2024e29d9 --- /dev/null +++ b/webview-ui/playwright/TranslationContext.ts @@ -0,0 +1,13 @@ +import { createContext, useContext } from "react" + +type TranslationContextValue = { + t: (key: string, options?: Record) => string + i18n: unknown +} + +export const TranslationContext = createContext({ + t: (key: string) => key, + i18n: null, +}) + +export const useAppTranslation = () => useContext(TranslationContext) diff --git a/webview-ui/src/components/settings/ModelDescriptionMarkdown.tsx b/webview-ui/src/components/settings/ModelDescriptionMarkdown.tsx index b04ab1163e..b00a36d1ac 100644 --- a/webview-ui/src/components/settings/ModelDescriptionMarkdown.tsx +++ b/webview-ui/src/components/settings/ModelDescriptionMarkdown.tsx @@ -3,7 +3,7 @@ import { memo, useEffect, useRef, useState } from "react" import { useRemark } from "react-remark" import { cn } from "@/lib/utils" -import { Collapsible, CollapsibleTrigger } from "@/components/ui" +import { Collapsible, CollapsibleTrigger } from "@/components/ui/collapsible" import { StyledMarkdown } from "./styles" diff --git a/webview-ui/src/components/settings/ModelInfoView.tsx b/webview-ui/src/components/settings/ModelInfoView.tsx index e043f68f83..fff55eda55 100644 --- a/webview-ui/src/components/settings/ModelInfoView.tsx +++ b/webview-ui/src/components/settings/ModelInfoView.tsx @@ -1,6 +1,7 @@ import { VSCodeLink } from "@vscode/webview-ui-toolkit/react" -import type { ModelInfo } from "@roo-code/types" +import { OpenAiServiceTier, type ModelInfo, type ServiceTier } from "@roo-code/types/model" +import { providerIdentifiers } from "@roo-code/types/provider-identifiers" import { formatPrice } from "@src/utils/formatPrice" import { cn } from "@src/lib/utils" @@ -17,6 +18,26 @@ type ModelInfoViewProps = { hidePricing?: boolean } +type TierPricingRowProps = { + tier: ServiceTier + label: string + modelInfo?: ModelInfo +} + +const TierPricingRow = ({ tier, label, modelInfo }: TierPricingRowProps) => { + const tierInfo = modelInfo?.tiers?.find(({ name }) => name === tier) + const fmt = (price?: number) => (typeof price === "number" ? formatPrice(price) : "—") + + return ( + + {label} + {fmt(tierInfo?.inputPrice ?? modelInfo?.inputPrice)} + {fmt(tierInfo?.outputPrice ?? modelInfo?.outputPrice)} + {fmt(tierInfo?.cacheReadsPrice ?? modelInfo?.cacheReadsPrice)} + + ) +} + export const ModelInfoView = ({ apiProvider, selectedModelId, @@ -29,9 +50,11 @@ export const ModelInfoView = ({ // Show tiered pricing table for OpenAI Native when model supports non-standard tiers const allowedTierNames = - modelInfo?.tiers?.filter((t) => t.name === "flex" || t.name === "priority")?.map((t) => t.name) ?? [] - const shouldShowTierPricingTable = apiProvider === "openai-native" && allowedTierNames.length > 0 - const fmt = (n?: number) => (typeof n === "number" ? `${formatPrice(n)}` : "—") + modelInfo?.tiers + ?.filter((t) => t.name === OpenAiServiceTier.Flex || t.name === OpenAiServiceTier.Priority) + ?.map((t) => t.name) ?? [] + const shouldShowTierPricingTable = apiProvider === providerIdentifiers.openaiNative && allowedTierNames.length > 0 + const fmt = (n?: number) => (typeof n === "number" ? formatPrice(n) : "—") const baseInfoItems = [ typeof modelInfo?.contextWindow === "number" && modelInfo.contextWindow > 0 && ( @@ -144,51 +167,19 @@ export const ModelInfoView = ({ {fmt(modelInfo?.outputPrice)} {fmt(modelInfo?.cacheReadsPrice)} - {allowedTierNames.includes("flex") && ( - - {t("settings:serviceTier.flex")} - - {fmt( - modelInfo?.tiers?.find((t) => t.name === "flex")?.inputPrice ?? - modelInfo?.inputPrice, - )} - - - {fmt( - modelInfo?.tiers?.find((t) => t.name === "flex")?.outputPrice ?? - modelInfo?.outputPrice, - )} - - - {fmt( - modelInfo?.tiers?.find((t) => t.name === "flex")?.cacheReadsPrice ?? - modelInfo?.cacheReadsPrice, - )} - - + {allowedTierNames.includes(OpenAiServiceTier.Flex) && ( + )} - {allowedTierNames.includes("priority") && ( - - {t("settings:serviceTier.priority")} - - {fmt( - modelInfo?.tiers?.find((t) => t.name === "priority")?.inputPrice ?? - modelInfo?.inputPrice, - )} - - - {fmt( - modelInfo?.tiers?.find((t) => t.name === "priority")?.outputPrice ?? - modelInfo?.outputPrice, - )} - - - {fmt( - modelInfo?.tiers?.find((t) => t.name === "priority")?.cacheReadsPrice ?? - modelInfo?.cacheReadsPrice, - )} - - + {allowedTierNames.includes(OpenAiServiceTier.Priority) && ( + )} diff --git a/webview-ui/src/components/settings/__tests__/ModelInfoView.spec.tsx b/webview-ui/src/components/settings/__tests__/ModelInfoView.spec.tsx new file mode 100644 index 0000000000..938bc8bff7 --- /dev/null +++ b/webview-ui/src/components/settings/__tests__/ModelInfoView.spec.tsx @@ -0,0 +1,129 @@ +import { OpenAiServiceTier, providerIdentifiers, type ModelInfo } from "@roo-code/types" + +import { render, screen, within } from "@/utils/test-utils" + +import { ModelInfoView } from "../ModelInfoView" + +vi.mock("@src/i18n/TranslationContext", () => ({ + useAppTranslation: () => ({ + t: (key: string) => + ({ + "settings:serviceTier.pricingTableTitle": "Service tier pricing", + "settings:serviceTier.columns.tier": "Tier", + "settings:serviceTier.columns.input": "Input", + "settings:serviceTier.columns.output": "Output", + "settings:serviceTier.columns.cacheReads": "Cache reads", + "settings:serviceTier.standard": "Standard", + "settings:serviceTier.flex": "Flex", + "settings:serviceTier.priority": "Priority", + })[key] ?? key, + }), +})) + +const baseModelInfo: ModelInfo = { + contextWindow: 128_000, + supportsPromptCache: true, + inputPrice: 10, + outputPrice: 20, + cacheReadsPrice: 3, +} + +const defaultProps = { + selectedModelId: "gpt-test", + isDescriptionExpanded: false, + setIsDescriptionExpanded: vi.fn(), +} + +const getPricingRowValues = (tier: string) => { + const row = screen.getByRole("cell", { name: tier }).closest("tr") + expect(row).not.toBeNull() + return within(row!) + .getAllByRole("cell") + .map((cell) => cell.textContent) +} + +describe("ModelInfoView service tier pricing", () => { + it("shows OpenAI Native tier prices with per-field fallback to Standard pricing", () => { + const modelInfo: ModelInfo = { + ...baseModelInfo, + tiers: [ + { name: OpenAiServiceTier.Default, contextWindow: 128_000 }, + { + name: OpenAiServiceTier.Flex, + contextWindow: 128_000, + inputPrice: 4, + cacheReadsPrice: 1, + }, + { + name: OpenAiServiceTier.Priority, + contextWindow: 128_000, + outputPrice: 40, + }, + ], + } + + render() + + expect(screen.getByText("Service tier pricing")).toBeInTheDocument() + expect(getPricingRowValues("Standard")).toEqual(["Standard", "$10.00", "$20.00", "$3.00"]) + expect(getPricingRowValues("Flex")).toEqual(["Flex", "$4.00", "$20.00", "$1.00"]) + expect(getPricingRowValues("Priority")).toEqual(["Priority", "$10.00", "$40.00", "$3.00"]) + }) + + it("only shows the tier pricing table for OpenAI Native models with a non-standard tier", () => { + const tieredModelInfo: ModelInfo = { + ...baseModelInfo, + tiers: [{ name: OpenAiServiceTier.Flex, contextWindow: 128_000 }], + } + const { rerender } = render( + , + ) + + expect(screen.queryByText("Service tier pricing")).not.toBeInTheDocument() + + rerender( + , + ) + + expect(screen.queryByText("Service tier pricing")).not.toBeInTheDocument() + }) + + it("does not show the tier pricing table when model info has no tiers", () => { + render( + , + ) + + expect(screen.queryByText("Service tier pricing")).not.toBeInTheDocument() + }) + + it("shows unavailable prices when neither the tier nor Standard defines them", () => { + render( + , + ) + + expect(getPricingRowValues("Flex")).toEqual(["Flex", "—", "—", "—"]) + expect(getPricingRowValues("Priority")).toEqual(["Priority", "—", "—", "—"]) + }) +}) diff --git a/webview-ui/src/components/settings/__tests__/ModelInfoView.visual.fixture.tsx b/webview-ui/src/components/settings/__tests__/ModelInfoView.visual.fixture.tsx new file mode 100644 index 0000000000..b007d4be69 --- /dev/null +++ b/webview-ui/src/components/settings/__tests__/ModelInfoView.visual.fixture.tsx @@ -0,0 +1,68 @@ +import React from "react" + +import { OpenAiServiceTier, type ModelInfo } from "@roo-code/types/model" +import { providerIdentifiers } from "@roo-code/types/provider-identifiers" + +import { TranslationContext } from "@src/i18n/TranslationContext" +import { ModelInfoView } from "../ModelInfoView" + +const modelInfo: ModelInfo = { + contextWindow: 400_000, + maxTokens: 128_000, + supportsImages: true, + supportsPromptCache: true, + inputPrice: 2, + outputPrice: 8, + cacheReadsPrice: 0.5, + tiers: [ + { + name: OpenAiServiceTier.Flex, + contextWindow: 400_000, + inputPrice: 1, + outputPrice: 4, + cacheReadsPrice: 0.25, + }, + { + name: OpenAiServiceTier.Priority, + contextWindow: 400_000, + inputPrice: 3.5, + outputPrice: 14, + cacheReadsPrice: 0.875, + }, + ], +} + +const translations: Record = { + "settings:modelInfo.contextWindow": "Context window:", + "settings:modelInfo.maxOutput": "Max output", + "settings:modelInfo.supportsImages": "Supports images", + "settings:modelInfo.noImages": "Does not support images", + "settings:modelInfo.supportsPromptCache": "Supports prompt caching", + "settings:modelInfo.noPromptCache": "Does not support prompt caching", + "settings:serviceTier.pricingTableTitle": "Service tier pricing (per 1M tokens)", + "settings:serviceTier.columns.tier": "Tier", + "settings:serviceTier.columns.input": "Input", + "settings:serviceTier.columns.output": "Output", + "settings:serviceTier.columns.cacheReads": "Cache reads", + "settings:serviceTier.standard": "Standard", + "settings:serviceTier.flex": "Flex", + "settings:serviceTier.priority": "Priority", +} + +export const ModelInfoViewFixture = () => ( + translations[key] ?? key, + i18n: null as unknown as typeof import("../../../i18n/setup").default, + }}> +
+ {}} + /> +
+
+) diff --git a/webview-ui/src/components/settings/__tests__/ModelInfoView.visual.tsx b/webview-ui/src/components/settings/__tests__/ModelInfoView.visual.tsx new file mode 100644 index 0000000000..0b5074e761 --- /dev/null +++ b/webview-ui/src/components/settings/__tests__/ModelInfoView.visual.tsx @@ -0,0 +1,15 @@ +import React from "react" + +import { expect, test } from "../../../../playwright/coverage-fixture" +import { ModelInfoViewFixture } from "./ModelInfoView.visual.fixture" + +test("renders OpenAI service tier pricing in the VS Code dark theme", async ({ mount }) => { + const component = await mount() + + await component.evaluate(async () => { + await document.fonts.ready + await new Promise((resolve) => requestAnimationFrame(() => resolve())) + }) + + await expect(component).toHaveScreenshot("model-info-service-tier-pricing-dark.png") +}) diff --git a/webview-ui/src/components/settings/__tests__/__screenshots__/model-info-service-tier-pricing-dark.png b/webview-ui/src/components/settings/__tests__/__screenshots__/model-info-service-tier-pricing-dark.png new file mode 100644 index 0000000000..0a3bee5351 Binary files /dev/null and b/webview-ui/src/components/settings/__tests__/__screenshots__/model-info-service-tier-pricing-dark.png differ diff --git a/webview-ui/src/components/settings/providers/OpenAI.tsx b/webview-ui/src/components/settings/providers/OpenAI.tsx index 96fd6c89be..460019e75a 100644 --- a/webview-ui/src/components/settings/providers/OpenAI.tsx +++ b/webview-ui/src/components/settings/providers/OpenAI.tsx @@ -2,7 +2,7 @@ import { useCallback, useState } from "react" import { Checkbox } from "vscrui" import { VSCodeTextField } from "@vscode/webview-ui-toolkit/react" -import type { ModelInfo, ProviderSettings } from "@roo-code/types" +import { OpenAiServiceTier, type ModelInfo, type ProviderSettings } from "@roo-code/types" import { useAppTranslation } from "@src/i18n/TranslationContext" import { VSCodeButtonLink } from "@src/components/common/VSCodeButtonLink" @@ -78,7 +78,7 @@ export const OpenAI = ({ apiConfiguration, setApiConfigurationField, selectedMod {(() => { const allowedTiers = (selectedModelInfo?.tiers?.map((t) => t.name).filter(Boolean) || []).filter( - (t) => t === "flex" || t === "priority", + (t) => t === OpenAiServiceTier.Flex || t === OpenAiServiceTier.Priority, ) if (allowedTiers.length === 0) return null @@ -92,7 +92,7 @@ export const OpenAI = ({ apiConfiguration, setApiConfigurationField, selectedMod diff --git a/webview-ui/src/components/settings/providers/__tests__/OpenAI.spec.tsx b/webview-ui/src/components/settings/providers/__tests__/OpenAI.spec.tsx new file mode 100644 index 0000000000..239c76932f --- /dev/null +++ b/webview-ui/src/components/settings/providers/__tests__/OpenAI.spec.tsx @@ -0,0 +1,92 @@ +import React from "react" + +import { OpenAiServiceTier, providerIdentifiers, type ModelInfo, type ProviderSettings } from "@roo-code/types" + +import { fireEvent, render, screen } from "@/utils/test-utils" + +import { OpenAI } from "../OpenAI" + +vi.mock("@src/i18n/TranslationContext", () => ({ + useAppTranslation: () => ({ t: (key: string) => key }), +})) + +vi.mock("vscrui", () => ({ + Checkbox: ({ children }: { children: React.ReactNode }) =>
{children}
, +})) + +vi.mock("@src/components/ui", () => ({ + Select: ({ children, value, onValueChange }: any) => ( + + ), + SelectContent: ({ children }: any) => <>{children}, + SelectItem: ({ children, value }: any) => , + SelectTrigger: () => null, + SelectValue: () => null, + StandardTooltip: ({ children, content }: any) => {children}, +})) + +const baseModelInfo: ModelInfo = { + contextWindow: 128_000, + supportsPromptCache: true, +} + +describe("OpenAI service tier selector", () => { + it("shows supported service tiers and persists the selected tier", () => { + const setApiConfigurationField = vi.fn() + const selectedModelInfo: ModelInfo = { + ...baseModelInfo, + tiers: [ + { name: OpenAiServiceTier.Default, contextWindow: 128_000 }, + { contextWindow: 128_000 }, + { name: OpenAiServiceTier.Flex, contextWindow: 128_000 }, + { name: OpenAiServiceTier.Priority, contextWindow: 128_000 }, + ], + } + const apiConfiguration: ProviderSettings = { + apiProvider: providerIdentifiers.openaiNative, + openAiNativeApiKey: "test-api-key", + } + + render( + , + ) + + const selector = screen.getByRole("combobox", { name: "Service tier" }) + expect(selector).toHaveValue(OpenAiServiceTier.Default) + expect(screen.getAllByRole("option").map((option) => option.textContent)).toEqual([ + "Standard", + "Flex", + "Priority", + ]) + + fireEvent.change(selector, { target: { value: OpenAiServiceTier.Flex } }) + expect(setApiConfigurationField).toHaveBeenLastCalledWith("openAiNativeServiceTier", OpenAiServiceTier.Flex) + + fireEvent.change(selector, { target: { value: OpenAiServiceTier.Priority } }) + expect(setApiConfigurationField).toHaveBeenLastCalledWith("openAiNativeServiceTier", OpenAiServiceTier.Priority) + }) + + it("hides the selector when the model only exposes the default tier", () => { + render( + , + ) + + expect(screen.queryByTestId("openai-service-tier")).not.toBeInTheDocument() + }) +}) diff --git a/webview-ui/src/components/welcome/__tests__/__screenshots__/zoo-hero-dark.png b/webview-ui/src/components/welcome/__tests__/__screenshots__/zoo-hero-dark.png index 6f5e11eaf9..3a99cfb487 100644 Binary files a/webview-ui/src/components/welcome/__tests__/__screenshots__/zoo-hero-dark.png and b/webview-ui/src/components/welcome/__tests__/__screenshots__/zoo-hero-dark.png differ diff --git a/webview-ui/src/index.css b/webview-ui/src/index.css index b93603f5a6..3f6cb9c54b 100644 --- a/webview-ui/src/index.css +++ b/webview-ui/src/index.css @@ -18,6 +18,8 @@ @import "tailwindcss/theme.css" layer(theme); @import "./preflight.css" layer(base); @import "tailwindcss/utilities.css" layer(utilities); + +@source "../src"; @import "katex/dist/katex.min.css"; @plugin "tailwindcss-animate";