diff --git a/src/api/providers/__tests__/bedrock.spec.ts b/src/api/providers/__tests__/bedrock.spec.ts index 0ea487eb44..975e38af12 100644 --- a/src/api/providers/__tests__/bedrock.spec.ts +++ b/src/api/providers/__tests__/bedrock.spec.ts @@ -1275,4 +1275,56 @@ describe("AwsBedrockHandler", () => { expect(mockCaptureException).toHaveBeenCalled() }) }) + + describe("prompt cache default behavior", () => { + beforeEach(() => { + mockConverseStreamCommand.mockReset() + }) + + // System prompt must exceed minTokensPerCachePoint (1024) for cache points to be placed + const longSystemPrompt = "You are a helpful assistant. ".repeat(200) + const messages: Anthropic.Messages.MessageParam[] = [{ role: "user", content: "Hello" }] + + it("should enable prompt caching by default when awsUsePromptCache is undefined", async () => { + const defaultHandler = new AwsBedrockHandler({ + apiModelId: "anthropic.claude-3-5-sonnet-20241022-v2:0", + awsAccessKey: "test-access-key", + awsSecretKey: "test-secret-key", + awsRegion: "us-east-1", + // awsUsePromptCache is intentionally omitted (undefined) + }) + + const generator = defaultHandler.createMessage(longSystemPrompt, messages) + await generator.next() // Start the generator + + expect(mockConverseStreamCommand).toHaveBeenCalled() + const commandArg = mockConverseStreamCommand.mock.calls[0][0] as any + + // System content should include a cachePoint entry since prompt caching defaults to ON + const systemBlocks = commandArg.system + const hasCachePoint = systemBlocks?.some((block: any) => block.cachePoint !== undefined) + expect(hasCachePoint).toBe(true) + }) + + it("should disable prompt caching when awsUsePromptCache is explicitly false", async () => { + const disabledHandler = new AwsBedrockHandler({ + apiModelId: "anthropic.claude-3-5-sonnet-20241022-v2:0", + awsAccessKey: "test-access-key", + awsSecretKey: "test-secret-key", + awsRegion: "us-east-1", + awsUsePromptCache: false, + }) + + const generator = disabledHandler.createMessage(longSystemPrompt, messages) + await generator.next() // Start the generator + + expect(mockConverseStreamCommand).toHaveBeenCalled() + const commandArg = mockConverseStreamCommand.mock.calls[0][0] as any + + // System content should NOT include cachePoint since caching is explicitly disabled + const systemBlocks = commandArg.system + const hasCachePoint = systemBlocks?.some((block: any) => block.cachePoint !== undefined) + expect(hasCachePoint).toBe(false) + }) + }) }) diff --git a/src/api/providers/bedrock.ts b/src/api/providers/bedrock.ts index 6bcf57d42a..3ceb251003 100644 --- a/src/api/providers/bedrock.ts +++ b/src/api/providers/bedrock.ts @@ -357,7 +357,9 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH }, ): ApiStream { const modelConfig = this.getModel() - const usePromptCache = Boolean(this.options.awsUsePromptCache && this.supportsAwsPromptCache(modelConfig)) + const usePromptCache = Boolean( + (this.options.awsUsePromptCache ?? true) && this.supportsAwsPromptCache(modelConfig), + ) const conversationId = messages.length > 0 diff --git a/webview-ui/src/components/settings/providers/Bedrock.tsx b/webview-ui/src/components/settings/providers/Bedrock.tsx index d9c69f8a8e..ed554f126d 100644 --- a/webview-ui/src/components/settings/providers/Bedrock.tsx +++ b/webview-ui/src/components/settings/providers/Bedrock.tsx @@ -198,7 +198,7 @@ export const Bedrock = ({ apiConfiguration, setApiConfigurationField, selectedMo {selectedModelInfo?.supportsPromptCache && ( <>
{t("settings:providers.enablePromptCaching")} diff --git a/webview-ui/src/components/ui/hooks/__tests__/useSelectedModel.spec.ts b/webview-ui/src/components/ui/hooks/__tests__/useSelectedModel.spec.ts index a8dead311f..2c24e4b565 100644 --- a/webview-ui/src/components/ui/hooks/__tests__/useSelectedModel.spec.ts +++ b/webview-ui/src/components/ui/hooks/__tests__/useSelectedModel.spec.ts @@ -496,6 +496,52 @@ describe("useSelectedModel", () => { }) }) + describe("bedrock provider with custom ARN", () => { + beforeEach(() => { + mockUseRouterModels.mockReturnValue({ + data: { + openrouter: {}, + requesty: {}, + litellm: {}, + }, + isLoading: false, + isError: false, + } as any) + + mockUseOpenRouterModelProviders.mockReturnValue({ + data: {}, + isLoading: false, + isError: false, + } as any) + }) + + it("should enable supportsPromptCache for custom-arn model", () => { + const apiConfiguration: ProviderSettings = { + apiProvider: "bedrock", + apiModelId: "custom-arn", + } + + const wrapper = createWrapper() + const { result } = renderHook(() => useSelectedModel(apiConfiguration), { wrapper }) + + expect(result.current.id).toBe("custom-arn") + expect(result.current.info?.supportsPromptCache).toBe(true) + }) + + it("should enable supportsImages for custom-arn model", () => { + const apiConfiguration: ProviderSettings = { + apiProvider: "bedrock", + apiModelId: "custom-arn", + } + + const wrapper = createWrapper() + const { result } = renderHook(() => useSelectedModel(apiConfiguration), { wrapper }) + + expect(result.current.id).toBe("custom-arn") + expect(result.current.info?.supportsImages).toBe(true) + }) + }) + describe("litellm provider", () => { beforeEach(() => { mockUseOpenRouterModelProviders.mockReturnValue({ diff --git a/webview-ui/src/components/ui/hooks/useSelectedModel.ts b/webview-ui/src/components/ui/hooks/useSelectedModel.ts index 8a6e49e212..959deff2b7 100644 --- a/webview-ui/src/components/ui/hooks/useSelectedModel.ts +++ b/webview-ui/src/components/ui/hooks/useSelectedModel.ts @@ -182,7 +182,7 @@ function getSelectedModel({ if (id === "custom-arn") { return { id, - info: { maxTokens: 5000, contextWindow: 128_000, supportsPromptCache: false, supportsImages: true }, + info: { maxTokens: 5000, contextWindow: 128_000, supportsPromptCache: true, supportsImages: true }, } }