enable aws bedrock provider with tool calling
This commit is contained in:
parent
1d748fb742
commit
ed272fb06d
7 changed files with 405 additions and 4 deletions
|
|
@ -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
18
go.mod
|
|
@ -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
36
go.sum
|
|
@ -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=
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
|
|
|
||||||
244
pkg/providers/aws_bedrock_provider.go
Normal file
244
pkg/providers/aws_bedrock_provider.go
Normal 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)
|
||||||
|
}
|
||||||
|
}
|
||||||
97
pkg/providers/aws_bedrock_provider_test.go
Normal file
97
pkg/providers/aws_bedrock_provider_test.go
Normal 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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
@ -430,4 +436,4 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) {
|
||||||
}
|
}
|
||||||
|
|
||||||
return NewHTTPProvider(apiKey, apiBase, proxy), nil
|
return NewHTTPProvider(apiKey, apiBase, proxy), nil
|
||||||
}
|
}
|
||||||
Loading…
Add table
Reference in a new issue