enable aws bedrock provider with tool calling

This commit is contained in:
Pranesh Shrestha 2026-02-15 22:42:29 +00:00 committed by Pranesh Shrestha
parent 1d748fb742
commit ed272fb06d
7 changed files with 405 additions and 4 deletions

View file

@ -107,6 +107,10 @@
"moonshot": { "moonshot": {
"api_key": "sk-xxx", "api_key": "sk-xxx",
"api_base": "" "api_base": ""
},
"aws_bedrock": {
"api_key": "",
"api_base": ""
} }
}, },
"tools": { "tools": {

18
go.mod
View file

@ -5,6 +5,9 @@ go 1.25.7
require ( require (
github.com/adhocore/gronx v1.19.6 github.com/adhocore/gronx v1.19.6
github.com/anthropics/anthropic-sdk-go v1.22.1 github.com/anthropics/anthropic-sdk-go v1.22.1
github.com/aws/aws-sdk-go-v2 v1.41.1
github.com/aws/aws-sdk-go-v2/config v1.32.7
github.com/aws/aws-sdk-go-v2/service/bedrockruntime v1.49.0
github.com/bwmarrin/discordgo v0.29.0 github.com/bwmarrin/discordgo v0.29.0
github.com/caarlos0/env/v11 v11.3.1 github.com/caarlos0/env/v11 v11.3.1
github.com/chzyer/readline v1.5.1 github.com/chzyer/readline v1.5.1
@ -28,12 +31,25 @@ require (
require ( require (
github.com/andybalholm/brotli v1.2.0 // indirect github.com/andybalholm/brotli v1.2.0 // indirect
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.4 // indirect
github.com/aws/aws-sdk-go-v2/credentials v1.19.7 // indirect
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.17 // indirect
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.17 // indirect
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.17 // indirect
github.com/aws/aws-sdk-go-v2/internal/ini v1.8.4 // indirect
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.4 // indirect
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.17 // indirect
github.com/aws/aws-sdk-go-v2/service/signin v1.0.5 // indirect
github.com/aws/aws-sdk-go-v2/service/sso v1.30.9 // indirect
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.35.13 // indirect
github.com/aws/aws-sdk-go-v2/service/sts v1.41.6 // indirect
github.com/aws/smithy-go v1.24.0 // indirect
github.com/bytedance/gopkg v0.1.3 // indirect github.com/bytedance/gopkg v0.1.3 // indirect
github.com/bytedance/sonic v1.15.0 // indirect github.com/bytedance/sonic v1.15.0 // indirect
github.com/bytedance/sonic/loader v0.5.0 // indirect github.com/bytedance/sonic/loader v0.5.0 // indirect
github.com/cloudwego/base64x v0.1.6 // indirect github.com/cloudwego/base64x v0.1.6 // indirect
github.com/github/copilot-sdk/go v0.1.23 github.com/github/copilot-sdk/go v0.1.23
github.com/go-resty/resty/v2 v2.17.1 // indirect github.com/go-resty/resty/v2 v2.17.2 // indirect
github.com/gogo/protobuf v1.3.2 // indirect github.com/gogo/protobuf v1.3.2 // indirect
github.com/google/jsonschema-go v0.4.2 // indirect github.com/google/jsonschema-go v0.4.2 // indirect
github.com/grbit/go-json v0.11.0 // indirect github.com/grbit/go-json v0.11.0 // indirect

36
go.sum
View file

@ -5,6 +5,38 @@ github.com/andybalholm/brotli v1.2.0 h1:ukwgCxwYrmACq68yiUqwIWnGY0cTPox/M94sVwTo
github.com/andybalholm/brotli v1.2.0/go.mod h1:rzTDkvFWvIrjDXZHkuS16NPggd91W3kUSvPlQ1pLaKY= github.com/andybalholm/brotli v1.2.0/go.mod h1:rzTDkvFWvIrjDXZHkuS16NPggd91W3kUSvPlQ1pLaKY=
github.com/anthropics/anthropic-sdk-go v1.22.1 h1:xbsc3vJKCX/ELDZSpTNfz9wCgrFsamwFewPb1iI0Xh0= github.com/anthropics/anthropic-sdk-go v1.22.1 h1:xbsc3vJKCX/ELDZSpTNfz9wCgrFsamwFewPb1iI0Xh0=
github.com/anthropics/anthropic-sdk-go v1.22.1/go.mod h1:WTz31rIUHUHqai2UslPpw5CwXrQP3geYBioRV4WOLvE= github.com/anthropics/anthropic-sdk-go v1.22.1/go.mod h1:WTz31rIUHUHqai2UslPpw5CwXrQP3geYBioRV4WOLvE=
github.com/aws/aws-sdk-go-v2 v1.41.1 h1:ABlyEARCDLN034NhxlRUSZr4l71mh+T5KAeGh6cerhU=
github.com/aws/aws-sdk-go-v2 v1.41.1/go.mod h1:MayyLB8y+buD9hZqkCW3kX1AKq07Y5pXxtgB+rRFhz0=
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.4 h1:489krEF9xIGkOaaX3CE/Be2uWjiXrkCH6gUX+bZA/BU=
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.4/go.mod h1:IOAPF6oT9KCsceNTvvYMNHy0+kMF8akOjeDvPENWxp4=
github.com/aws/aws-sdk-go-v2/config v1.32.7 h1:vxUyWGUwmkQ2g19n7JY/9YL8MfAIl7bTesIUykECXmY=
github.com/aws/aws-sdk-go-v2/config v1.32.7/go.mod h1:2/Qm5vKUU/r7Y+zUk/Ptt2MDAEKAfUtKc1+3U1Mo3oY=
github.com/aws/aws-sdk-go-v2/credentials v1.19.7 h1:tHK47VqqtJxOymRrNtUXN5SP/zUTvZKeLx4tH6PGQc8=
github.com/aws/aws-sdk-go-v2/credentials v1.19.7/go.mod h1:qOZk8sPDrxhf+4Wf4oT2urYJrYt3RejHSzgAquYeppw=
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.17 h1:I0GyV8wiYrP8XpA70g1HBcQO1JlQxCMTW9npl5UbDHY=
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.17/go.mod h1:tyw7BOl5bBe/oqvoIeECFJjMdzXoa/dfVz3QQ5lgHGA=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.17 h1:xOLELNKGp2vsiteLsvLPwxC+mYmO6OZ8PYgiuPJzF8U=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.17/go.mod h1:5M5CI3D12dNOtH3/mk6minaRwI2/37ifCURZISxA/IQ=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.17 h1:WWLqlh79iO48yLkj1v3ISRNiv+3KdQoZ6JWyfcsyQik=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.17/go.mod h1:EhG22vHRrvF8oXSTYStZhJc1aUgKtnJe+aOiFEV90cM=
github.com/aws/aws-sdk-go-v2/internal/ini v1.8.4 h1:WKuaxf++XKWlHWu9ECbMlha8WOEGm0OUEZqm4K/Gcfk=
github.com/aws/aws-sdk-go-v2/internal/ini v1.8.4/go.mod h1:ZWy7j6v1vWGmPReu0iSGvRiise4YI5SkR3OHKTZ6Wuc=
github.com/aws/aws-sdk-go-v2/service/bedrockruntime v1.49.0 h1:osqN479arsxXAIHmBbiAn+0nj7jCkuXtzgtZPSwt0sc=
github.com/aws/aws-sdk-go-v2/service/bedrockruntime v1.49.0/go.mod h1:siKVmJdui4dwPPtsKr3F5BAeJxW1MANWaLJnTDfgu7c=
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.4 h1:0ryTNEdJbzUCEWkVXEXoqlXV72J5keC1GvILMOuD00E=
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.4/go.mod h1:HQ4qwNZh32C3CBeO6iJLQlgtMzqeG17ziAA/3KDJFow=
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.17 h1:RuNSMoozM8oXlgLG/n6WLaFGoea7/CddrCfIiSA+xdY=
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.17/go.mod h1:F2xxQ9TZz5gDWsclCtPQscGpP0VUOc8RqgFM3vDENmU=
github.com/aws/aws-sdk-go-v2/service/signin v1.0.5 h1:VrhDvQib/i0lxvr3zqlUwLwJP4fpmpyD9wYG1vfSu+Y=
github.com/aws/aws-sdk-go-v2/service/signin v1.0.5/go.mod h1:k029+U8SY30/3/ras4G/Fnv/b88N4mAfliNn08Dem4M=
github.com/aws/aws-sdk-go-v2/service/sso v1.30.9 h1:v6EiMvhEYBoHABfbGB4alOYmCIrcgyPPiBE1wZAEbqk=
github.com/aws/aws-sdk-go-v2/service/sso v1.30.9/go.mod h1:yifAsgBxgJWn3ggx70A3urX2AN49Y5sJTD1UQFlfqBw=
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.35.13 h1:gd84Omyu9JLriJVCbGApcLzVR3XtmC4ZDPcAI6Ftvds=
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.35.13/go.mod h1:sTGThjphYE4Ohw8vJiRStAcu3rbjtXRsdNB0TvZ5wwo=
github.com/aws/aws-sdk-go-v2/service/sts v1.41.6 h1:5fFjR/ToSOzB2OQ/XqWpZBmNvmP/pJ1jOWYlFDJTjRQ=
github.com/aws/aws-sdk-go-v2/service/sts v1.41.6/go.mod h1:qgFDZQSD/Kys7nJnVqYlWKnh0SSdMjAi0uSwON4wgYQ=
github.com/aws/smithy-go v1.24.0 h1:LpilSUItNPFr1eY85RYgTIg5eIEPtvFbskaFcmmIUnk=
github.com/aws/smithy-go v1.24.0/go.mod h1:LEj2LM3rBRQJxPZTB4KuzZkaZYnZPnvgIhb4pu07mx0=
github.com/bwmarrin/discordgo v0.29.0 h1:FmWeXFaKUwrcL3Cx65c20bTRW+vOb6k8AnaP+EgjDno= github.com/bwmarrin/discordgo v0.29.0 h1:FmWeXFaKUwrcL3Cx65c20bTRW+vOb6k8AnaP+EgjDno=
github.com/bwmarrin/discordgo v0.29.0/go.mod h1:NJZpH+1AfhIcyQsPeuBKsUtYrRnjkyu0kIVMCHkZtRY= github.com/bwmarrin/discordgo v0.29.0/go.mod h1:NJZpH+1AfhIcyQsPeuBKsUtYrRnjkyu0kIVMCHkZtRY=
github.com/bytedance/gopkg v0.1.3 h1:TPBSwH8RsouGCBcMBktLt1AymVo2TVsBVCY4b6TnZ/M= github.com/bytedance/gopkg v0.1.3 h1:TPBSwH8RsouGCBcMBktLt1AymVo2TVsBVCY4b6TnZ/M=
@ -36,8 +68,8 @@ github.com/github/copilot-sdk/go v0.1.23 h1:uExtO/inZQndCZMiSAA1hvXINiz9tqo/MZgQ
github.com/github/copilot-sdk/go v0.1.23/go.mod h1:GdwwBfMbm9AABLEM3x5IZKw4ZfwCYxZ1BgyytmZenQ0= github.com/github/copilot-sdk/go v0.1.23/go.mod h1:GdwwBfMbm9AABLEM3x5IZKw4ZfwCYxZ1BgyytmZenQ0=
github.com/go-redis/redis/v8 v8.11.4/go.mod h1:2Z2wHZXdQpCDXEGzqMockDpNyYvi2l4Pxt6RJr792+w= github.com/go-redis/redis/v8 v8.11.4/go.mod h1:2Z2wHZXdQpCDXEGzqMockDpNyYvi2l4Pxt6RJr792+w=
github.com/go-resty/resty/v2 v2.6.0/go.mod h1:PwvJS6hvaPkjtjNg9ph+VrSD92bi5Zq73w/BIH7cC3Q= github.com/go-resty/resty/v2 v2.6.0/go.mod h1:PwvJS6hvaPkjtjNg9ph+VrSD92bi5Zq73w/BIH7cC3Q=
github.com/go-resty/resty/v2 v2.17.1 h1:x3aMpHK1YM9e4va/TMDRlusDDoZiQ+ViDu/WpA6xTM4= github.com/go-resty/resty/v2 v2.17.2 h1:FQW5oHYcIlkCNrMD2lloGScxcHJ0gkjshV3qcQAyHQk=
github.com/go-resty/resty/v2 v2.17.1/go.mod h1:kCKZ3wWmwJaNc7S29BRtUhJwy7iqmn+2mLtQrOyQlVA= github.com/go-resty/resty/v2 v2.17.2/go.mod h1:kCKZ3wWmwJaNc7S29BRtUhJwy7iqmn+2mLtQrOyQlVA=
github.com/go-task/slim-sprig v0.0.0-20210107165309-348f09dbbbc0/go.mod h1:fyg7847qk6SyHyPtNmDHnmrv/HOrqktSC+C9fM+CJOE= github.com/go-task/slim-sprig v0.0.0-20210107165309-348f09dbbbc0/go.mod h1:fyg7847qk6SyHyPtNmDHnmrv/HOrqktSC+C9fM+CJOE=
github.com/go-test/deep v1.1.1 h1:0r/53hagsehfO4bzD2Pgr/+RgHqhmf+k1Bpse2cTu1U= github.com/go-test/deep v1.1.1 h1:0r/53hagsehfO4bzD2Pgr/+RgHqhmf+k1Bpse2cTu1U=
github.com/go-test/deep v1.1.1/go.mod h1:5C2ZWiW0ErCdrYzpqxLbTX7MG14M9iiw8DgHncVwcsE= github.com/go-test/deep v1.1.1/go.mod h1:5C2ZWiW0ErCdrYzpqxLbTX7MG14M9iiw8DgHncVwcsE=

View file

@ -178,6 +178,7 @@ type ProvidersConfig struct {
Moonshot ProviderConfig `json:"moonshot"` Moonshot ProviderConfig `json:"moonshot"`
ShengSuanYun ProviderConfig `json:"shengsuanyun"` ShengSuanYun ProviderConfig `json:"shengsuanyun"`
DeepSeek ProviderConfig `json:"deepseek"` DeepSeek ProviderConfig `json:"deepseek"`
AWSBedrock ProviderConfig `json:"awsbedrock"`
GitHubCopilot ProviderConfig `json:"github_copilot"` GitHubCopilot ProviderConfig `json:"github_copilot"`
} }
@ -304,6 +305,7 @@ func DefaultConfig() *Config {
Nvidia: ProviderConfig{}, Nvidia: ProviderConfig{},
Moonshot: ProviderConfig{}, Moonshot: ProviderConfig{},
ShengSuanYun: ProviderConfig{}, ShengSuanYun: ProviderConfig{},
AWSBedrock: ProviderConfig{},
}, },
Gateway: GatewayConfig{ Gateway: GatewayConfig{
Host: "0.0.0.0", Host: "0.0.0.0",

View file

@ -0,0 +1,244 @@
package providers
import (
"context"
"encoding/json"
"fmt"
"log"
"strings"
"github.com/aws/aws-sdk-go-v2/aws"
"github.com/aws/aws-sdk-go-v2/config"
"github.com/aws/aws-sdk-go-v2/service/bedrockruntime"
"github.com/aws/aws-sdk-go-v2/service/bedrockruntime/document"
"github.com/aws/aws-sdk-go-v2/service/bedrockruntime/types"
)
type AWSBedrockProvider struct {
command string
workspace string
}
func NewAWSBedrockProvider(workspace string) *AWSBedrockProvider {
return &AWSBedrockProvider{
command: "awsbedrock",
workspace: workspace,
}
}
func (p *AWSBedrockProvider) Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error) {
BedrockRuntimeClient, err := getBedrockRuntimeClient()
if err != nil {
return nil, err
}
converseInput := p.messagesToConverseInput(messages, model, tools)
response, err := BedrockRuntimeClient.Converse(ctx, &converseInput)
if err != nil {
processError(err, model)
return nil, err
}
return p.parseAWSBedrockResponse(response)
}
func (p *AWSBedrockProvider) GetDefaultModel() string {
return "anthropic.claude-haiku-4-5-20251001-v1:0"
}
func (p *AWSBedrockProvider) messagesToConverseInput(messages []Message, model string, tools []ToolDefinition) bedrockruntime.ConverseInput {
var systemBlocks []types.SystemContentBlock
var conversationMessages []types.Message
for _, msg := range messages {
if msg.Role == "system" {
systemBlocks = append(systemBlocks, &types.SystemContentBlockMemberText{Value: msg.Content})
continue
}
var bedrockRole types.ConversationRole
var contentBlocks []types.ContentBlock
switch msg.Role {
case "user":
bedrockRole = types.ConversationRoleUser
if msg.ToolCallID != "" {
contentBlocks = append(contentBlocks, &types.ContentBlockMemberToolResult{
Value: types.ToolResultBlock{
ToolUseId: aws.String(msg.ToolCallID),
Content: []types.ToolResultContentBlock{
&types.ToolResultContentBlockMemberText{Value: msg.Content},
},
},
})
} else {
contentBlocks = append(contentBlocks, &types.ContentBlockMemberText{Value: msg.Content})
}
case "assistant":
bedrockRole = types.ConversationRoleAssistant
if len(msg.ToolCalls) > 0 {
if msg.Content != "" {
contentBlocks = append(contentBlocks, &types.ContentBlockMemberText{Value: msg.Content})
}
for _, tc := range msg.ToolCalls {
name := tc.Name
if name == "" && tc.Function != nil {
name = tc.Function.Name
}
args := tc.Arguments
if len(args) == 0 && tc.Function != nil && tc.Function.Arguments != "" {
var parsed map[string]interface{}
if err := json.Unmarshal([]byte(tc.Function.Arguments), &parsed); err == nil {
args = parsed
}
}
contentBlocks = append(contentBlocks, &types.ContentBlockMemberToolUse{
Value: types.ToolUseBlock{
ToolUseId: aws.String(tc.ID),
Name: aws.String(name),
Input: document.NewLazyDocument(args),
},
})
}
} else {
contentBlocks = append(contentBlocks, &types.ContentBlockMemberText{Value: msg.Content})
}
case "tool":
bedrockRole = types.ConversationRoleUser
contentBlocks = append(contentBlocks, &types.ContentBlockMemberToolResult{
Value: types.ToolResultBlock{
ToolUseId: aws.String(msg.ToolCallID),
Content: []types.ToolResultContentBlock{
&types.ToolResultContentBlockMemberText{Value: msg.Content},
},
},
})
default:
continue
}
if len(conversationMessages) > 0 && conversationMessages[len(conversationMessages)-1].Role == bedrockRole {
lastMsg := &conversationMessages[len(conversationMessages)-1]
lastMsg.Content = append(lastMsg.Content, contentBlocks...)
} else {
conversationMessages = append(conversationMessages, types.Message{
Role: bedrockRole,
Content: contentBlocks,
})
}
}
input := bedrockruntime.ConverseInput{
ModelId: aws.String(model),
Messages: conversationMessages,
}
if len(systemBlocks) > 0 {
input.System = systemBlocks
}
if len(tools) > 0 {
var toolConfigs []types.Tool
for _, tool := range tools {
toolConfigs = append(toolConfigs, &types.ToolMemberToolSpec{
Value: types.ToolSpecification{
Name: aws.String(tool.Function.Name),
Description: aws.String(tool.Function.Description),
InputSchema: &types.ToolInputSchemaMemberJson{
Value: document.NewLazyDocument(tool.Function.Parameters),
},
},
})
}
input.ToolConfig = &types.ToolConfiguration{
Tools: toolConfigs,
}
}
return input
}
// parseAWSBedrockResponse parses the JSON output from the AWS Bedrock API.
func (p *AWSBedrockProvider) parseAWSBedrockResponse(response *bedrockruntime.ConverseOutput) (*LLMResponse, error) {
outputMsg, ok := response.Output.(*types.ConverseOutputMemberMessage)
if !ok {
return nil, fmt.Errorf("unexpected output type")
}
message := outputMsg.Value
var content strings.Builder
var toolCalls []ToolCall
for _, block := range message.Content {
switch b := block.(type) {
case *types.ContentBlockMemberText:
if content.Len() > 0 {
content.WriteString("\n")
}
content.WriteString(b.Value)
case *types.ContentBlockMemberToolUse:
toolUse := b.Value
args := map[string]interface{}{}
if toolUse.Input != nil {
if inputBytes, err := toolUse.Input.MarshalSmithyDocument(); err == nil {
json.Unmarshal(inputBytes, &args)
}
}
toolCalls = append(toolCalls, ToolCall{
ID: *toolUse.ToolUseId,
Name: *toolUse.Name,
Arguments: args,
})
}
}
finishReason := "stop"
if response.StopReason != "" {
switch response.StopReason {
case types.StopReasonToolUse:
finishReason = "tool_calls"
case types.StopReasonMaxTokens:
finishReason = "length"
case types.StopReasonEndTurn:
finishReason = "stop"
}
}
var usage *UsageInfo
if response.Usage != nil {
usage = &UsageInfo{
PromptTokens: int(*response.Usage.InputTokens),
CompletionTokens: int(*response.Usage.OutputTokens),
TotalTokens: int(*response.Usage.TotalTokens),
}
}
return &LLMResponse{
Content: content.String(),
ToolCalls: toolCalls,
FinishReason: finishReason,
Usage: usage,
}, nil
}
func getBedrockRuntimeClient() (*bedrockruntime.Client, error) {
cfg, err := config.LoadDefaultConfig(context.TODO())
if err != nil {
log.Fatalf("failed to load config: %v", err)
}
client := bedrockruntime.NewFromConfig(cfg)
return client, nil
}
func processError(err error, modelId string) {
errMsg := err.Error()
if strings.Contains(errMsg, "no such host") {
fmt.Printf(`The Bedrock service is not available in the selected region.
Please double-check the service availability for your region at
https://aws.amazon.com/about-aws/global-infrastructure/regional-product-services/.\n`)
} else if strings.Contains(errMsg, "Could not resolve the foundation model") {
fmt.Printf(`Could not resolve the foundation model from model identifier: \"%v\".
Please verify that the requested model exists and is accessible
within the specified region.\n
`, modelId)
} else {
fmt.Printf("Couldn't invoke model: \"%v\". Here's why: %v\n", modelId, err)
}
}

View file

@ -0,0 +1,97 @@
package providers
import (
"testing"
"github.com/aws/aws-sdk-go-v2/service/bedrockruntime/types"
)
func TestMessagesToConverseInput(t *testing.T) {
p := NewAWSBedrockProvider(".")
tests := []struct {
name string
messages []Message
expectedSystemCount int
expectedMessageCount int
expectedRoles []types.ConversationRole
}{
{
name: "single system and user message",
messages: []Message{
{Role: "system", Content: "System prompt"},
{Role: "user", Content: "Hello"},
},
expectedSystemCount: 1,
expectedMessageCount: 1,
expectedRoles: []types.ConversationRole{types.ConversationRoleUser},
},
{
name: "alternating roles",
messages: []Message{
{Role: "user", Content: "Hello"},
{Role: "assistant", Content: "Hi there"},
{Role: "user", Content: "How are you?"},
},
expectedSystemCount: 0,
expectedMessageCount: 3,
expectedRoles: []types.ConversationRole{types.ConversationRoleUser, types.ConversationRoleAssistant, types.ConversationRoleUser},
},
{
name: "consecutive user messages should merge",
messages: []Message{
{Role: "user", Content: "Message 1"},
{Role: "user", Content: "Message 2"},
},
expectedSystemCount: 0,
expectedMessageCount: 1,
expectedRoles: []types.ConversationRole{types.ConversationRoleUser},
},
{
name: "user message then tool result should merge",
messages: []Message{
{Role: "user", Content: "Run tool"},
{Role: "tool", Content: "Result", ToolCallID: "call_1"},
},
expectedSystemCount: 0,
expectedMessageCount: 1,
expectedRoles: []types.ConversationRole{types.ConversationRoleUser},
},
{
name: "complex sequence with merging",
messages: []Message{
{Role: "system", Content: "Sys 1"},
{Role: "system", Content: "Sys 2"},
{Role: "user", Content: "User 1"},
{Role: "assistant", Content: "Assist 1"},
{Role: "assistant", Content: "Assist 2"},
{Role: "user", Content: "User 2"},
{Role: "tool", Content: "Res 1", ToolCallID: "call_a"},
{Role: "tool", Content: "Res 2", ToolCallID: "call_b"},
},
expectedSystemCount: 2,
expectedMessageCount: 3, // User 1, Assistant 1+2, User 2+Res 1+Res 2
expectedRoles: []types.ConversationRole{types.ConversationRoleUser, types.ConversationRoleAssistant, types.ConversationRoleUser},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
input := p.messagesToConverseInput(tt.messages, "some-model", nil)
if len(input.System) != tt.expectedSystemCount {
t.Errorf("expected %d system blocks, got %d", tt.expectedSystemCount, len(input.System))
}
if len(input.Messages) != tt.expectedMessageCount {
t.Errorf("expected %d messages, got %d", tt.expectedMessageCount, len(input.Messages))
}
for i, role := range tt.expectedRoles {
if i < len(input.Messages) && input.Messages[i].Role != role {
t.Errorf("expected role %s at index %d, got %s", role, i, input.Messages[i].Role)
}
}
})
}
}

View file

@ -323,6 +323,12 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) {
} }
return NewGitHubCopilotProvider(apiBase, cfg.Providers.GitHubCopilot.ConnectMode, model) return NewGitHubCopilotProvider(apiBase, cfg.Providers.GitHubCopilot.ConnectMode, model)
case "awsbedrock":
workspace := cfg.Agents.Defaults.Workspace
if workspace == "" {
workspace = "."
}
return NewAWSBedrockProvider(workspace), nil
} }
} }