Implement extended reasoning support in AwsBedrockHandler and add corresponding tests

This commit is contained in:
hannesrudolph 2025-06-07 12:23:59 -06:00
parent 395f55b31f
commit 3d794b9f71
4 changed files with 454 additions and 17 deletions

View file

@ -73,6 +73,7 @@ export const bedrockModels = {
supportsImages: true,
supportsComputerUse: true,
supportsPromptCache: true,
supportsReasoningBudget: true,
inputPrice: 3.0,
outputPrice: 15.0,
cacheWritesPrice: 3.75,
@ -87,6 +88,7 @@ export const bedrockModels = {
supportsImages: true,
supportsComputerUse: true,
supportsPromptCache: true,
supportsReasoningBudget: true,
inputPrice: 15.0,
outputPrice: 75.0,
cacheWritesPrice: 18.75,
@ -101,6 +103,7 @@ export const bedrockModels = {
supportsImages: true,
supportsComputerUse: true,
supportsPromptCache: true,
supportsReasoningBudget: true,
inputPrice: 3.0,
outputPrice: 15.0,
cacheWritesPrice: 3.75,

View file

@ -0,0 +1,248 @@
import { AwsBedrockHandler } from "../bedrock"
import { BedrockRuntimeClient } from "@aws-sdk/client-bedrock-runtime"
// Mock the AWS SDK
jest.mock("@aws-sdk/client-bedrock-runtime")
jest.mock("@aws-sdk/credential-providers")
describe("AwsBedrockHandler - Extended Thinking/Reasoning", () => {
let mockClient: jest.Mocked<BedrockRuntimeClient>
let mockSend: jest.Mock
beforeEach(() => {
jest.clearAllMocks()
mockSend = jest.fn()
mockClient = {
send: mockSend,
config: { region: "us-east-1" },
} as any
;(BedrockRuntimeClient as jest.Mock).mockImplementation(() => mockClient)
})
describe("Extended Thinking Configuration", () => {
it("should NOT include thinking configuration by default", async () => {
const handler = new AwsBedrockHandler({
apiModelId: "anthropic.claude-3-7-sonnet-20250219-v1:0",
awsAccessKey: "test-key",
awsSecretKey: "test-secret",
awsRegion: "us-east-1",
// enableReasoningEffort is NOT set, so reasoning should be disabled
})
// Mock the stream response
mockSend.mockResolvedValueOnce({
stream: (async function* () {
yield { messageStart: { role: "assistant" } }
yield { contentBlockStart: { start: { text: "Hello" } } }
yield { messageStop: {} }
})(),
})
const messages = [{ role: "user" as const, content: "Test message" }]
const stream = handler.createMessage("System prompt", messages)
// Consume the stream
for await (const _chunk of stream) {
// Just consume
}
// Verify the command was called
expect(mockSend).toHaveBeenCalledTimes(1)
const command = mockSend.mock.calls[0][0]
const payload = command.input
// Verify thinking is NOT included
expect(payload.anthropic_version).toBeUndefined()
expect(payload.additionalModelRequestFields).toBeUndefined()
expect(payload.inferenceConfig.temperature).toBeDefined()
expect(payload.inferenceConfig.topP).toBeDefined()
})
it("should include thinking configuration when explicitly enabled", async () => {
const handler = new AwsBedrockHandler({
apiModelId: "anthropic.claude-3-7-sonnet-20250219-v1:0",
awsAccessKey: "test-key",
awsSecretKey: "test-secret",
awsRegion: "us-east-1",
enableReasoningEffort: true, // Explicitly enable reasoning
modelMaxThinkingTokens: 5000, // Set thinking tokens
})
// Mock the stream response
mockSend.mockResolvedValueOnce({
stream: (async function* () {
yield { messageStart: { role: "assistant" } }
yield { contentBlockStart: { contentBlock: { type: "thinking", thinking: "Let me think..." } } }
yield { contentBlockDelta: { delta: { type: "thinking_delta", thinking: " about this." } } }
yield { contentBlockStart: { start: { text: "Here's my answer" } } }
yield { messageStop: {} }
})(),
})
const messages = [{ role: "user" as const, content: "Test message" }]
const stream = handler.createMessage("System prompt", messages)
// Consume the stream
const chunks = []
for await (const chunk of stream) {
chunks.push(chunk)
}
// Verify the command was called
expect(mockSend).toHaveBeenCalledTimes(1)
const command = mockSend.mock.calls[0][0]
const payload = command.input
// Verify thinking IS included
expect(payload.anthropic_version).toBe("bedrock-20250514")
expect(payload.additionalModelRequestFields).toEqual({
thinking: {
type: "enabled",
budget_tokens: 5000,
},
})
// Temperature and topP should be removed when thinking is enabled
expect(payload.inferenceConfig.temperature).toBeUndefined()
expect(payload.inferenceConfig.topP).toBeUndefined()
// Verify thinking chunks were properly handled
const reasoningChunks = chunks.filter((c) => c.type === "reasoning")
expect(reasoningChunks).toHaveLength(2)
expect(reasoningChunks[0].text).toBe("Let me think...")
expect(reasoningChunks[1].text).toBe(" about this.")
})
it("should NOT enable thinking for non-supported models even if requested", async () => {
const handler = new AwsBedrockHandler({
apiModelId: "anthropic.claude-3-haiku-20240307-v1:0", // This model doesn't support reasoning
awsAccessKey: "test-key",
awsSecretKey: "test-secret",
awsRegion: "us-east-1",
enableReasoningEffort: true, // Try to enable reasoning
modelMaxThinkingTokens: 5000,
})
// Mock the stream response
mockSend.mockResolvedValueOnce({
stream: (async function* () {
yield { messageStart: { role: "assistant" } }
yield { contentBlockStart: { start: { text: "Hello" } } }
yield { messageStop: {} }
})(),
})
const messages = [{ role: "user" as const, content: "Test message" }]
const stream = handler.createMessage("System prompt", messages)
// Consume the stream
for await (const _chunk of stream) {
// Just consume
}
// Verify the command was called
expect(mockSend).toHaveBeenCalledTimes(1)
const command = mockSend.mock.calls[0][0]
const payload = command.input
// Verify thinking is NOT included because model doesn't support it
expect(payload.anthropic_version).toBeUndefined()
expect(payload.additionalModelRequestFields).toBeUndefined()
})
it("should handle thinking stream events correctly", async () => {
const handler = new AwsBedrockHandler({
apiModelId: "anthropic.claude-sonnet-4-20250514-v1:0",
awsAccessKey: "test-key",
awsSecretKey: "test-secret",
awsRegion: "us-east-1",
enableReasoningEffort: true,
modelMaxThinkingTokens: 8000,
})
// Mock the stream response with various thinking events
mockSend.mockResolvedValueOnce({
stream: (async function* () {
yield { messageStart: { role: "assistant" } }
// Thinking block start
yield {
contentBlockStart: { contentBlock: { type: "thinking", thinking: "Analyzing the request..." } },
}
// Thinking deltas
yield {
contentBlockDelta: { delta: { type: "thinking_delta", thinking: "\nThis seems complex." } },
}
yield {
contentBlockDelta: { delta: { type: "thinking_delta", thinking: "\nLet me break it down." } },
}
// Signature delta (part of thinking)
yield {
contentBlockDelta: { delta: { type: "signature_delta", signature: "\n[Signature: ABC123]" } },
}
// Regular text response
yield { contentBlockStart: { start: { text: "Based on my analysis" } } }
yield { contentBlockDelta: { delta: { text: ", here's the answer." } } }
yield { messageStop: {} }
})(),
})
const messages = [{ role: "user" as const, content: "Complex question" }]
const stream = handler.createMessage("System prompt", messages)
// Collect all chunks
const chunks = []
for await (const chunk of stream) {
chunks.push(chunk)
}
// Verify reasoning chunks
const reasoningChunks = chunks.filter((c) => c.type === "reasoning")
expect(reasoningChunks).toHaveLength(4)
expect(reasoningChunks.map((c) => c.text).join("")).toBe(
"Analyzing the request...\nThis seems complex.\nLet me break it down.\n[Signature: ABC123]",
)
// Verify text chunks
const textChunks = chunks.filter((c) => c.type === "text")
expect(textChunks).toHaveLength(2)
expect(textChunks.map((c) => c.text).join("")).toBe("Based on my analysis, here's the answer.")
})
})
describe("Error Handling for Extended Thinking", () => {
it("should provide helpful error message for thinking-related errors", async () => {
const handler = new AwsBedrockHandler({
apiModelId: "anthropic.claude-3-7-sonnet-20250219-v1:0",
awsAccessKey: "test-key",
awsSecretKey: "test-secret",
awsRegion: "us-east-1",
enableReasoningEffort: true,
modelMaxThinkingTokens: 5000,
})
// Mock an error response
const error = new Error("ValidationException: additionalModelRequestFields.thinking is not supported")
mockSend.mockRejectedValueOnce(error)
const messages = [{ role: "user" as const, content: "Test message" }]
const stream = handler.createMessage("System prompt", messages)
// Collect error chunks
const chunks = []
try {
for await (const chunk of stream) {
chunks.push(chunk)
}
} catch (e) {
// Expected to throw
}
// Should have error chunks before throwing
expect(chunks).toHaveLength(2)
expect(chunks[0].type).toBe("text")
if (chunks[0].type === "text") {
expect(chunks[0].text).toContain("Extended thinking/reasoning is not supported")
}
expect(chunks[1].type).toBe("usage")
})
})
})

View file

@ -30,6 +30,8 @@ import { MultiPointStrategy } from "../transform/cache-strategy/multi-point-stra
import { ModelInfo as CacheModelInfo } from "../transform/cache-strategy/types"
import { convertToBedrockConverseMessages as sharedConverter } from "../transform/bedrock-converse-format"
import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata } from "../index"
import { getModelParams } from "../transform/model-params"
import { shouldUseReasoningBudget } from "../../shared/api"
/************************************************************************************
*
@ -40,8 +42,8 @@ import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata } from ".
// Define interface for Bedrock inference config
interface BedrockInferenceConfig {
maxTokens: number
temperature: number
topP: number
temperature?: number
topP?: number
}
// Define types for stream events based on AWS SDK
@ -58,10 +60,20 @@ export interface StreamEvent {
text?: string
}
contentBlockIndex?: number
// Extended thinking support
contentBlock?: {
type?: "thinking" | "text"
thinking?: string
text?: string
}
}
contentBlockDelta?: {
delta?: {
text?: string
// Extended thinking support
type?: "thinking_delta" | "text_delta" | "signature_delta"
thinking?: string
signature?: string
}
contentBlockIndex?: number
}
@ -256,6 +268,49 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH
systemPrompt: string,
messages: Anthropic.Messages.MessageParam[],
metadata?: ApiHandlerCreateMessageMetadata,
): ApiStream {
const maxRetries = 3
let retryCount = 0
let lastError: unknown
while (retryCount < maxRetries) {
try {
yield* this.createMessageInternal(systemPrompt, messages, metadata)
return
} catch (error) {
lastError = error
retryCount++
// Check if error is retryable
const errorType = this.getErrorType(error)
const retryableErrors = ["THROTTLING", "ABORT", "GENERIC"]
if (!retryableErrors.includes(errorType) || retryCount >= maxRetries) {
// Not retryable or max retries reached
throw error
}
// Log retry attempt
logger.info(`Retrying Bedrock request (attempt ${retryCount}/${maxRetries})`, {
ctx: "bedrock",
errorType,
errorMessage: error instanceof Error ? error.message : String(error),
})
// Exponential backoff: 1s, 2s, 4s
const delay = Math.pow(2, retryCount - 1) * 1000
await new Promise((resolve) => setTimeout(resolve, delay))
}
}
// If we get here, all retries failed
throw lastError
}
private async *createMessageInternal(
systemPrompt: string,
messages: Anthropic.Messages.MessageParam[],
metadata?: ApiHandlerCreateMessageMetadata,
): ApiStream {
let modelConfig = this.getModel()
// Handle cross-region inference
@ -280,20 +335,59 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH
conversationId,
)
// Get model parameters including reasoning configuration
const params = getModelParams({
format: "anthropic",
modelId: modelConfig.id as string,
model: modelConfig.info,
settings: this.options,
})
// Construct the payload
const inferenceConfig: BedrockInferenceConfig = {
maxTokens: modelConfig.info.maxTokens as number,
temperature: this.options.modelTemperature as number,
maxTokens: params.maxTokens || (modelConfig.info.maxTokens as number),
temperature: params.temperature || (this.options.modelTemperature as number),
topP: 0.1,
}
const payload = {
// Build the base payload
const payload: any = {
modelId: modelConfig.id,
messages: formatted.messages,
system: formatted.system,
inferenceConfig,
}
// Add extended thinking support ONLY if explicitly enabled by the user
// Reasoning is disabled by default as per requirements
if (
this.options.enableReasoningEffort &&
shouldUseReasoningBudget({ model: modelConfig.info, settings: this.options }) &&
params.reasoning &&
params.reasoningBudget
) {
// Add the anthropic_version field required for extended thinking
payload.anthropic_version = "bedrock-20250514"
// Add additionalModelRequestFields with thinking configuration
payload.additionalModelRequestFields = {
thinking: {
type: "enabled",
budget_tokens: params.reasoningBudget,
},
}
// Remove temperature, topP, and top_k when thinking is enabled as they are incompatible
delete inferenceConfig.temperature
delete inferenceConfig.topP
logger.info("Extended thinking enabled for Bedrock request", {
ctx: "bedrock",
modelId: modelConfig.id,
budgetTokens: params.reasoningBudget,
})
}
// Create AbortController with 10 minute timeout
const controller = new AbortController()
let timeoutId: NodeJS.Timeout | undefined
@ -397,21 +491,59 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH
}
// Handle content blocks
if (streamEvent.contentBlockStart?.start?.text) {
yield {
type: "text",
text: streamEvent.contentBlockStart.start.text,
if (streamEvent.contentBlockStart) {
// Handle thinking content blocks
if (streamEvent.contentBlockStart.contentBlock?.type === "thinking") {
yield {
type: "reasoning",
text: streamEvent.contentBlockStart.contentBlock.thinking || "",
}
continue
}
// Handle regular text content blocks
if (streamEvent.contentBlockStart.start?.text || streamEvent.contentBlockStart.contentBlock?.text) {
yield {
type: "text",
text:
streamEvent.contentBlockStart.start?.text ||
streamEvent.contentBlockStart.contentBlock?.text ||
"",
}
continue
}
continue
}
// Handle content deltas
if (streamEvent.contentBlockDelta?.delta?.text) {
yield {
type: "text",
text: streamEvent.contentBlockDelta.delta.text,
if (streamEvent.contentBlockDelta?.delta) {
const delta = streamEvent.contentBlockDelta.delta
// Handle thinking deltas
if (delta.type === "thinking_delta" && delta.thinking) {
yield {
type: "reasoning",
text: delta.thinking,
}
continue
}
// Handle signature deltas (part of thinking)
if (delta.type === "signature_delta" && delta.signature) {
// Signature is part of the thinking process, treat it as reasoning
yield {
type: "reasoning",
text: delta.signature,
}
continue
}
// Handle regular text deltas
if (delta.text) {
yield {
type: "text",
text: delta.text,
}
continue
}
continue
}
// Handle message stop
if (streamEvent.messageStop) {
@ -509,7 +641,14 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH
conversationId?: string, // Optional conversation ID to track cache points across messages
): { system: SystemContentBlock[]; messages: Message[] } {
// First convert messages using shared converter for proper image handling
const convertedMessages = sharedConverter(anthropicMessages as Anthropic.Messages.MessageParam[])
let convertedMessages = sharedConverter(anthropicMessages as Anthropic.Messages.MessageParam[])
// Handle extended thinking for tool use
// When using tools with extended thinking, we need to preserve the thinking block
// from previous assistant messages
if (this.options.enableReasoningEffort && modelInfo?.supportsReasoningBudget) {
convertedMessages = this.preserveThinkingBlocks(convertedMessages)
}
// If prompt caching is disabled, return the converted messages directly
if (!usePromptCache) {
@ -792,6 +931,35 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH
return content
}
/**
* Preserves thinking blocks from previous assistant messages for tool use continuity
*/
private preserveThinkingBlocks(messages: Message[]): Message[] {
// When using extended thinking with tools, we need to preserve the entire
// thinking block from previous assistant messages to maintain reasoning continuity
return messages.map((message, index) => {
if (message.role === "assistant" && index > 0) {
// Check if this assistant message follows a tool use pattern
const prevMessage = messages[index - 1]
if (prevMessage.role === "user" && this.hasToolUseContent(prevMessage)) {
// This is likely a response to tool use, preserve any thinking blocks
return message
}
}
return message
})
}
/**
* Checks if a message contains tool use content
*/
private hasToolUseContent(message: Message): boolean {
if (!message.content || !Array.isArray(message.content)) {
return false
}
return message.content.some((block: any) => block.toolUse || block.toolResult)
}
/************************************************************************************
*
* AMAZON REGIONS
@ -905,10 +1073,22 @@ Suggestions:
messageTemplate: `Invalid ARN format. ARN should follow the pattern: arn:aws:bedrock:region:account-id:resource-type/resource-name`,
logLevel: "error",
},
THINKING_NOT_SUPPORTED: {
patterns: ["thinking", "reasoning", "additionalmodelrequestfields"],
messageTemplate: `Extended thinking/reasoning is not supported for this model or configuration.
Please verify:
1. You're using a supported model (Claude 3.7 Sonnet, Claude 4 Sonnet, or Claude 4 Opus)
2. Your AWS region supports extended thinking
3. You have the necessary permissions to use this feature
If the issue persists, try disabling "Enable Reasoning Mode" in the settings.`,
logLevel: "error",
},
// Default/generic error
GENERIC: {
patterns: [], // Empty patterns array means this is the default
messageTemplate: `Unknown Error`,
messageTemplate: `Bedrock is unable to process your request. Please check your configuration and try again.`,
logLevel: "error",
},
}

View file

@ -8,6 +8,7 @@ import { useAppTranslation } from "@src/i18n/TranslationContext"
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@src/components/ui"
import { inputEventTransform, noTransform } from "../transforms"
import { ThinkingBudget } from "../ThinkingBudget"
type BedrockProps = {
apiConfiguration: ProviderSettings
@ -151,6 +152,11 @@ export const Bedrock = ({ apiConfiguration, setApiConfigurationField, selectedMo
</div>
</>
)}
<ThinkingBudget
apiConfiguration={apiConfiguration}
setApiConfigurationField={setApiConfigurationField}
modelInfo={selectedModelInfo}
/>
</>
)
}