diff --git a/packages/types/src/provider-settings.ts b/packages/types/src/provider-settings.ts index 5262e7602d..e1af2126e5 100644 --- a/packages/types/src/provider-settings.ts +++ b/packages/types/src/provider-settings.ts @@ -219,6 +219,7 @@ const bedrockSchema = apiModelIdProviderModelSchema.extend({ awsSessionToken: z.string().optional(), awsRegion: z.string().optional(), awsUseCrossRegionInference: z.boolean().optional(), + awsUseGlobalInference: z.boolean().optional(), awsUsePromptCache: z.boolean().optional(), awsProfile: z.string().optional(), awsUseProfile: z.boolean().optional(), diff --git a/packages/types/src/providers/bedrock.ts b/packages/types/src/providers/bedrock.ts index 251757cad8..c97465f2b3 100644 --- a/packages/types/src/providers/bedrock.ts +++ b/packages/types/src/providers/bedrock.ts @@ -442,6 +442,21 @@ export const AWS_INFERENCE_PROFILE_MAPPING: Array<[string, string]> = [ ["sa-", "sa."], ] +// Global Inference Profile prefix +// https://docs.aws.amazon.com/bedrock/latest/userguide/global-inference.html +export const AWS_GLOBAL_INFERENCE_PREFIX = "global." + +// Models that support Global Inference +// Based on AWS documentation, these models can use the global. prefix +export const BEDROCK_GLOBAL_INFERENCE_MODEL_IDS = [ + "anthropic.claude-sonnet-4-20250514-v1:0", + "anthropic.claude-sonnet-4-5-20250929-v1:0", + "anthropic.claude-opus-4-20250514-v1:0", + "anthropic.claude-opus-4-1-20250805-v1:0", + "anthropic.claude-3-7-sonnet-20250219-v1:0", + "anthropic.claude-haiku-4-5-20251001-v1:0", +] as const + // Amazon Bedrock supported regions for the regions dropdown // Based on official AWS documentation export const BEDROCK_REGIONS = [ diff --git a/src/api/providers/__tests__/bedrock-global-inference.spec.ts b/src/api/providers/__tests__/bedrock-global-inference.spec.ts new file mode 100644 index 0000000000..168b3504d2 --- /dev/null +++ b/src/api/providers/__tests__/bedrock-global-inference.spec.ts @@ -0,0 +1,272 @@ +// npx vitest run src/api/providers/__tests__/bedrock-global-inference.spec.ts + +import { AWS_GLOBAL_INFERENCE_PREFIX, BEDROCK_GLOBAL_INFERENCE_MODEL_IDS } from "@roo-code/types" +import { AwsBedrockHandler } from "../bedrock" +import { ApiHandlerOptions } from "../../../shared/api" + +// Mock AWS SDK +vitest.mock("@aws-sdk/client-bedrock-runtime", () => { + return { + BedrockRuntimeClient: vitest.fn().mockImplementation(() => ({ + send: vitest.fn(), + config: { region: "us-east-1" }, + })), + ConverseCommand: vitest.fn(), + ConverseStreamCommand: vitest.fn(), + } +}) + +describe("Amazon Bedrock Global Inference", () => { + // Helper function to create a handler with specific options + const createHandler = (options: Partial = {}) => { + const defaultOptions: ApiHandlerOptions = { + apiModelId: "anthropic.claude-sonnet-4-20250514-v1:0", + awsRegion: "us-east-1", + ...options, + } + return new AwsBedrockHandler(defaultOptions) + } + + describe("AWS_GLOBAL_INFERENCE_PREFIX constant", () => { + it("should have the correct global inference prefix", () => { + expect(AWS_GLOBAL_INFERENCE_PREFIX).toBe("global.") + }) + }) + + describe("BEDROCK_GLOBAL_INFERENCE_MODEL_IDS constant", () => { + it("should contain the expected models that support global inference", () => { + const expectedModels = [ + "anthropic.claude-sonnet-4-20250514-v1:0", + "anthropic.claude-sonnet-4-5-20250929-v1:0", + "anthropic.claude-opus-4-20250514-v1:0", + "anthropic.claude-opus-4-1-20250805-v1:0", + "anthropic.claude-3-7-sonnet-20250219-v1:0", + "anthropic.claude-haiku-4-5-20251001-v1:0", + ] + expect(BEDROCK_GLOBAL_INFERENCE_MODEL_IDS).toEqual(expectedModels) + }) + }) + + describe("Global inference with supported models", () => { + it("should apply global. prefix when global inference is enabled for Claude Sonnet 4", () => { + const handler = createHandler({ + awsUseGlobalInference: true, + apiModelId: "anthropic.claude-sonnet-4-20250514-v1:0", + }) + + const model = handler.getModel() + expect(model.id).toBe("global.anthropic.claude-sonnet-4-20250514-v1:0") + }) + + it("should apply global. prefix when global inference is enabled for Claude Sonnet 4.5", () => { + const handler = createHandler({ + awsUseGlobalInference: true, + apiModelId: "anthropic.claude-sonnet-4-5-20250929-v1:0", + }) + + const model = handler.getModel() + expect(model.id).toBe("global.anthropic.claude-sonnet-4-5-20250929-v1:0") + }) + + it("should apply global. prefix when global inference is enabled for Claude Opus 4", () => { + const handler = createHandler({ + awsUseGlobalInference: true, + apiModelId: "anthropic.claude-opus-4-20250514-v1:0", + }) + + const model = handler.getModel() + expect(model.id).toBe("global.anthropic.claude-opus-4-20250514-v1:0") + }) + + it("should apply global. prefix when global inference is enabled for Claude Opus 4.1", () => { + const handler = createHandler({ + awsUseGlobalInference: true, + apiModelId: "anthropic.claude-opus-4-1-20250805-v1:0", + }) + + const model = handler.getModel() + expect(model.id).toBe("global.anthropic.claude-opus-4-1-20250805-v1:0") + }) + + it("should apply global. prefix when global inference is enabled for Claude 3.7 Sonnet", () => { + const handler = createHandler({ + awsUseGlobalInference: true, + apiModelId: "anthropic.claude-3-7-sonnet-20250219-v1:0", + }) + + const model = handler.getModel() + expect(model.id).toBe("global.anthropic.claude-3-7-sonnet-20250219-v1:0") + }) + + it("should apply global. prefix when global inference is enabled for Claude Haiku 4.5", () => { + const handler = createHandler({ + awsUseGlobalInference: true, + apiModelId: "anthropic.claude-haiku-4-5-20251001-v1:0", + }) + + const model = handler.getModel() + expect(model.id).toBe("global.anthropic.claude-haiku-4-5-20251001-v1:0") + }) + }) + + describe("Global inference with unsupported models", () => { + it("should NOT apply global. prefix for unsupported Claude 3 Sonnet model", () => { + const handler = createHandler({ + awsUseGlobalInference: true, + apiModelId: "anthropic.claude-3-sonnet-20240229-v1:0", + }) + + const model = handler.getModel() + expect(model.id).toBe("anthropic.claude-3-sonnet-20240229-v1:0") + }) + + it("should NOT apply global. prefix for unsupported Claude 3 Haiku model", () => { + const handler = createHandler({ + awsUseGlobalInference: true, + apiModelId: "anthropic.claude-3-haiku-20240307-v1:0", + }) + + const model = handler.getModel() + expect(model.id).toBe("anthropic.claude-3-haiku-20240307-v1:0") + }) + + it("should NOT apply global. prefix for Amazon Nova models", () => { + const handler = createHandler({ + awsUseGlobalInference: true, + apiModelId: "amazon.nova-pro-v1:0", + }) + + const model = handler.getModel() + expect(model.id).toBe("amazon.nova-pro-v1:0") + }) + + it("should NOT apply global. prefix for Llama models", () => { + const handler = createHandler({ + awsUseGlobalInference: true, + apiModelId: "meta.llama3-1-70b-instruct-v1:0", + }) + + const model = handler.getModel() + expect(model.id).toBe("meta.llama3-1-70b-instruct-v1:0") + }) + }) + + describe("Global inference priority over cross-region inference", () => { + it("should prioritize global inference over cross-region inference when both are enabled", () => { + const handler = createHandler({ + awsUseGlobalInference: true, + awsUseCrossRegionInference: true, + awsRegion: "us-east-1", + apiModelId: "anthropic.claude-sonnet-4-20250514-v1:0", + }) + + const model = handler.getModel() + // Should use global. prefix, not us. prefix + expect(model.id).toBe("global.anthropic.claude-sonnet-4-20250514-v1:0") + }) + + it("should fall back to cross-region inference when global is disabled", () => { + const handler = createHandler({ + awsUseGlobalInference: false, + awsUseCrossRegionInference: true, + awsRegion: "us-east-1", + apiModelId: "anthropic.claude-sonnet-4-20250514-v1:0", + }) + + const model = handler.getModel() + // Should use us. prefix for cross-region inference + expect(model.id).toBe("us.anthropic.claude-sonnet-4-20250514-v1:0") + }) + + it("should apply no prefix when both global and cross-region are disabled", () => { + const handler = createHandler({ + awsUseGlobalInference: false, + awsUseCrossRegionInference: false, + awsRegion: "us-east-1", + apiModelId: "anthropic.claude-sonnet-4-20250514-v1:0", + }) + + const model = handler.getModel() + // Should have no prefix + expect(model.id).toBe("anthropic.claude-sonnet-4-20250514-v1:0") + }) + }) + + describe("Global inference with custom ARNs", () => { + it("should parse global inference from ARN", () => { + const handler = createHandler({ + awsCustomArn: + "arn:aws:bedrock:us-east-1:123456789012:inference-profile/global.anthropic.claude-sonnet-4-20250514-v1:0", + }) + + const model = handler.getModel() + expect(model.id).toBe("global.anthropic.claude-sonnet-4-20250514-v1:0") + }) + + it("should distinguish between global and cross-region prefixes in ARNs", () => { + // Test global inference ARN + const globalHandler = createHandler({ + awsCustomArn: + "arn:aws:bedrock:us-east-1:123456789012:inference-profile/global.anthropic.claude-sonnet-4-20250514-v1:0", + }) + const globalModel = globalHandler.getModel() + expect(globalModel.id).toBe("global.anthropic.claude-sonnet-4-20250514-v1:0") + + // Test cross-region inference ARN + const crossRegionHandler = createHandler({ + awsCustomArn: + "arn:aws:bedrock:us-east-1:123456789012:inference-profile/us.anthropic.claude-3-sonnet-20240229-v1:0", + }) + const crossRegionModel = crossRegionHandler.getModel() + expect(crossRegionModel.id).toBe("us.anthropic.claude-3-sonnet-20240229-v1:0") + }) + }) + + describe("parseBaseModelId function", () => { + it("should remove global. prefix from model IDs", () => { + const handler = createHandler() + const parseBaseModelId = (handler as any).parseBaseModelId.bind(handler) + + expect(parseBaseModelId("global.anthropic.claude-sonnet-4-20250514-v1:0")).toBe( + "anthropic.claude-sonnet-4-20250514-v1:0", + ) + expect(parseBaseModelId("global.anthropic.claude-opus-4-20250514-v1:0")).toBe( + "anthropic.claude-opus-4-20250514-v1:0", + ) + }) + + it("should remove cross-region prefixes from model IDs", () => { + const handler = createHandler() + const parseBaseModelId = (handler as any).parseBaseModelId.bind(handler) + + expect(parseBaseModelId("us.anthropic.claude-3-sonnet-20240229-v1:0")).toBe( + "anthropic.claude-3-sonnet-20240229-v1:0", + ) + expect(parseBaseModelId("eu.anthropic.claude-3-sonnet-20240229-v1:0")).toBe( + "anthropic.claude-3-sonnet-20240229-v1:0", + ) + expect(parseBaseModelId("apac.anthropic.claude-3-sonnet-20240229-v1:0")).toBe( + "anthropic.claude-3-sonnet-20240229-v1:0", + ) + }) + + it("should prioritize global. prefix removal over cross-region prefixes", () => { + const handler = createHandler() + const parseBaseModelId = (handler as any).parseBaseModelId.bind(handler) + + // Even if there's a model ID that somehow has both (shouldn't happen in practice), + // global. should be removed first + expect(parseBaseModelId("global.us.some-model-id")).toBe("us.some-model-id") + }) + + it("should return model ID unchanged if no prefix is present", () => { + const handler = createHandler() + const parseBaseModelId = (handler as any).parseBaseModelId.bind(handler) + + expect(parseBaseModelId("anthropic.claude-3-sonnet-20240229-v1:0")).toBe( + "anthropic.claude-3-sonnet-20240229-v1:0", + ) + expect(parseBaseModelId("amazon.nova-pro-v1:0")).toBe("amazon.nova-pro-v1:0") + }) + }) +}) diff --git a/src/api/providers/bedrock.ts b/src/api/providers/bedrock.ts index 9267fb924b..4c98cb7eab 100644 --- a/src/api/providers/bedrock.ts +++ b/src/api/providers/bedrock.ts @@ -22,6 +22,8 @@ import { BEDROCK_DEFAULT_CONTEXT, AWS_INFERENCE_PROFILE_MAPPING, BEDROCK_1M_CONTEXT_MODEL_IDS, + AWS_GLOBAL_INFERENCE_PREFIX, + BEDROCK_GLOBAL_INFERENCE_MODEL_IDS, } from "@roo-code/types" import { ApiStream } from "../transform/stream" @@ -209,7 +211,8 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH } this.options.apiModelId = this.arnInfo.modelId - if (this.arnInfo.awsUseCrossRegionInference) this.options.awsUseCrossRegionInference = true + if (this.arnInfo.crossRegionInference) this.options.awsUseCrossRegionInference = true + if (this.arnInfo.globalInference) this.options.awsUseGlobalInference = true } if (!this.options.modelTemperature) { @@ -832,9 +835,11 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH modelId?: string errorMessage?: string crossRegionInference: boolean + globalInference?: boolean } = { isValid: true, crossRegionInference: false, // Default to false + globalInference: false, // Default to false } result.modelType = match[3] @@ -845,11 +850,17 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH const arnRegion = match[1] result.region = arnRegion - // Check if the original model ID had a region prefix + // Check if the original model ID had a prefix if (originalModelId && result.modelId !== originalModelId) { - // If the model ID changed after parsing, it had a region prefix + // If the model ID changed after parsing, it had a prefix let prefix = originalModelId.replace(result.modelId, "") - result.crossRegionInference = AwsBedrockHandler.isSystemInferenceProfile(prefix) + + // Check if it's a global inference prefix + if (prefix === AWS_GLOBAL_INFERENCE_PREFIX) { + result.globalInference = true + } else { + result.crossRegionInference = AwsBedrockHandler.isSystemInferenceProfile(prefix) + } } // Check if region in ARN matches provided region (if specified) @@ -878,6 +889,11 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH return modelId } + // Remove AWS global inference prefix first + if (modelId.startsWith(AWS_GLOBAL_INFERENCE_PREFIX)) { + return modelId.substring(AWS_GLOBAL_INFERENCE_PREFIX.length) + } + // Remove AWS cross-region inference profile prefixes // as defined in AWS_INFERENCE_PROFILE_MAPPING for (const [_, inferenceProfile] of AWS_INFERENCE_PROFILE_MAPPING) { @@ -958,17 +974,44 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH //If the user entered an ARN for a foundation-model they've done the same thing as picking from our list of options. //We leave the model data matching the same as if a drop-down input method was used by not overwriting the model ID with the user input ARN - //Otherwise the ARN is not a foundation-model resource type that ARN should be used as the identifier in Bedrock interactions - if (this.arnInfo.modelType !== "foundation-model") modelConfig.id = this.options.awsCustomArn + //For inference-profile ARNs with global or cross-region prefixes, reconstruct the model ID with the prefix + //Otherwise the ARN should be used as the identifier in Bedrock interactions + if (this.arnInfo.modelType === "inference-profile") { + if (this.arnInfo.globalInference) { + // Re-add the global. prefix that was stripped during parsing + modelConfig.id = `${AWS_GLOBAL_INFERENCE_PREFIX}${this.arnInfo.modelId}` + } else if (this.arnInfo.crossRegionInference) { + // For cross-region, we need to determine the prefix based on the region + const prefix = AwsBedrockHandler.getPrefixForRegion(this.options.awsRegion!) + if (prefix) { + modelConfig.id = `${prefix}${this.arnInfo.modelId}` + } else { + // Fallback to the original ARN if we can't determine the prefix + modelConfig.id = this.options.awsCustomArn + } + } else { + // No special prefix, use the ARN as-is + modelConfig.id = this.options.awsCustomArn + } + } else if (this.arnInfo.modelType !== "foundation-model") { + modelConfig.id = this.options.awsCustomArn + } } else { //a model was selected from the drop down modelConfig = this.getModelById(this.options.apiModelId as string) - // Add cross-region inference prefix if enabled - if (this.options.awsUseCrossRegionInference && this.options.awsRegion) { + // Use parseBaseModelId to get the clean model ID for checking against support lists + const baseModelId = this.parseBaseModelId(modelConfig.id) + + // Add global inference prefix if enabled and model supports it + if (this.options.awsUseGlobalInference && BEDROCK_GLOBAL_INFERENCE_MODEL_IDS.includes(baseModelId as any)) { + modelConfig.id = `${AWS_GLOBAL_INFERENCE_PREFIX}${baseModelId}` + } + // Add cross-region inference prefix if enabled (and global inference is not being used) + else if (this.options.awsUseCrossRegionInference && this.options.awsRegion) { const prefix = AwsBedrockHandler.getPrefixForRegion(this.options.awsRegion) if (prefix) { - modelConfig.id = `${prefix}${modelConfig.id}` + modelConfig.id = `${prefix}${baseModelId}` } } } diff --git a/webview-ui/src/components/settings/providers/Bedrock.tsx b/webview-ui/src/components/settings/providers/Bedrock.tsx index 1b3143fa08..29feb2e248 100644 --- a/webview-ui/src/components/settings/providers/Bedrock.tsx +++ b/webview-ui/src/components/settings/providers/Bedrock.tsx @@ -2,7 +2,13 @@ import { useCallback, useState, useEffect } from "react" import { Checkbox } from "vscrui" import { VSCodeTextField } from "@vscode/webview-ui-toolkit/react" -import { type ProviderSettings, type ModelInfo, BEDROCK_REGIONS, BEDROCK_1M_CONTEXT_MODEL_IDS } from "@roo-code/types" +import { + type ProviderSettings, + type ModelInfo, + BEDROCK_REGIONS, + BEDROCK_1M_CONTEXT_MODEL_IDS, + BEDROCK_GLOBAL_INFERENCE_MODEL_IDS, +} from "@roo-code/types" import { useAppTranslation } from "@src/i18n/TranslationContext" import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue, StandardTooltip } from "@src/components/ui" @@ -23,6 +29,11 @@ export const Bedrock = ({ apiConfiguration, setApiConfigurationField, selectedMo const supports1MContextBeta = !!apiConfiguration?.apiModelId && BEDROCK_1M_CONTEXT_MODEL_IDS.includes(apiConfiguration.apiModelId as any) + // Check if the selected model supports global inference + const supportsGlobalInference = + !!apiConfiguration?.apiModelId && + BEDROCK_GLOBAL_INFERENCE_MODEL_IDS.includes(apiConfiguration.apiModelId as any) + // Update the endpoint enabled state when the configuration changes useEffect(() => { setAwsEndpointSelected(!!apiConfiguration?.awsBedrockEndpointEnabled) @@ -138,9 +149,17 @@ export const Bedrock = ({ apiConfiguration, setApiConfigurationField, selectedMo + {supportsGlobalInference && ( + + {t("settings:providers.awsGlobalInference")} + + )} + onChange={handleInputChange("awsUseCrossRegionInference", noTransform)} + disabled={apiConfiguration?.awsUseGlobalInference}> {t("settings:providers.awsCrossRegion")} {selectedModelInfo?.supportsPromptCache && ( diff --git a/webview-ui/src/i18n/locales/en/settings.json b/webview-ui/src/i18n/locales/en/settings.json index dfccc49cc4..0e08bb0799 100644 --- a/webview-ui/src/i18n/locales/en/settings.json +++ b/webview-ui/src/i18n/locales/en/settings.json @@ -341,6 +341,7 @@ "awsSecretKey": "AWS Secret Key", "awsSessionToken": "AWS Session Token", "awsRegion": "AWS Region", + "awsGlobalInference": "Use global inference", "awsCrossRegion": "Use cross-region inference", "awsBedrockVpc": { "useCustomVpcEndpoint": "Use custom VPC endpoint",