feat: add global inference support for AWS Bedrock models

- Add awsUseGlobalInference setting to provider settings schema
- Add support for global. prefix in Bedrock provider for supported models
- Add UI checkbox for global inference in Bedrock settings
- Prioritize global inference over cross-region when both are enabled
- Add comprehensive tests for global inference functionality
- Support global inference in custom ARN parsing

Supported models:
- Claude Sonnet 4, 4.5
- Claude Opus 4, 4.1
- Claude 3.7 Sonnet
- Claude Haiku 4.5

Fixes #8750
This commit is contained in:
Roo Code 2025-10-21 17:07:29 +00:00
parent f3a505fd36
commit 9b4c3264a5
6 changed files with 362 additions and 11 deletions

View file

@ -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(),

View file

@ -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 = [

View file

@ -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<ApiHandlerOptions> = {}) => {
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")
})
})
})

View file

@ -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}`
}
}
}

View file

@ -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
</SelectContent>
</Select>
</div>
{supportsGlobalInference && (
<Checkbox
checked={apiConfiguration?.awsUseGlobalInference || false}
onChange={handleInputChange("awsUseGlobalInference", noTransform)}>
{t("settings:providers.awsGlobalInference")}
</Checkbox>
)}
<Checkbox
checked={apiConfiguration?.awsUseCrossRegionInference || false}
onChange={handleInputChange("awsUseCrossRegionInference", noTransform)}>
onChange={handleInputChange("awsUseCrossRegionInference", noTransform)}
disabled={apiConfiguration?.awsUseGlobalInference}>
{t("settings:providers.awsCrossRegion")}
</Checkbox>
{selectedModelInfo?.supportsPromptCache && (

View file

@ -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",