Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 5 additions & 4 deletions go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,9 @@ require (
github.com/GoCodeAlone/workflow v0.85.2
github.com/GoCodeAlone/workflow-plugin-authz v0.5.4
github.com/anthropics/anthropic-sdk-go v1.26.0
github.com/aws/aws-sdk-go-v2 v1.41.6
github.com/aws/aws-sdk-go-v2 v1.42.1
github.com/aws/aws-sdk-go-v2/credentials v1.19.15
github.com/aws/aws-sdk-go-v2/service/bedrock v1.64.2
github.com/bmatcuk/doublestar/v4 v4.10.0
github.com/coder/acp-go-sdk v0.6.3
github.com/creack/pty v1.1.24
Expand Down Expand Up @@ -56,8 +57,8 @@ require (
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.8 // indirect
github.com/aws/aws-sdk-go-v2/config v1.32.16 // indirect
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.22 // indirect
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.22 // indirect
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.22 // indirect
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.30 // indirect
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.30 // indirect
github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.23 // indirect
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.8 // indirect
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.22 // indirect
Expand All @@ -66,7 +67,7 @@ require (
github.com/aws/aws-sdk-go-v2/service/sso v1.30.16 // indirect
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.35.20 // indirect
github.com/aws/aws-sdk-go-v2/service/sts v1.42.0 // indirect
github.com/aws/smithy-go v1.25.0 // indirect
github.com/aws/smithy-go v1.27.3 // indirect
github.com/bahlo/generic-list-go v0.2.0 // indirect
github.com/beorn7/perks v1.0.1 // indirect
github.com/bits-and-blooms/bitset v1.24.5 // indirect
Expand Down
18 changes: 10 additions & 8 deletions go.sum
Original file line number Diff line number Diff line change
Expand Up @@ -64,8 +64,8 @@ github.com/antithesishq/antithesis-sdk-go v0.7.2 h1:oEEedg1Xgi8drRjqB0f9tfjhLoIn
github.com/antithesishq/antithesis-sdk-go v0.7.2/go.mod h1:FQyySiasQQM8735Ddel3MRojmy4dA1IqCeyJ5jmPMbI=
github.com/armon/go-metrics v0.4.1 h1:hR91U9KYmb6bLBYLQjyM+3j+rcd/UhE+G78SFnF8gJA=
github.com/armon/go-metrics v0.4.1/go.mod h1:E6amYzXo6aW1tqzoZGT755KkbgrJsSdpwZ+3JqfkOG4=
github.com/aws/aws-sdk-go-v2 v1.41.6 h1:1AX0AthnBQzMx1vbmir3Y4WsnJgiydmnJjiLu+LvXOg=
github.com/aws/aws-sdk-go-v2 v1.41.6/go.mod h1:dy0UzBIfwSeot4grGvY1AqFWN5zgziMmWGzysDnHFcQ=
github.com/aws/aws-sdk-go-v2 v1.42.1 h1:9eOTgu1z/dVtYpNZ3/8/XbbaX0x/BqE3HUzAzs6K0ek=
github.com/aws/aws-sdk-go-v2 v1.42.1/go.mod h1:5pKeft2eJj+gElQ38Jqg4ibCqh+/AK33/0X3hip7IjM=
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.8 h1:eBMB84YGghSocM7PsjmmPffTa+1FBUeNvGvFou6V/4o=
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.8/go.mod h1:lyw7GFp3qENLh7kwzf7iMzAxDn+NzjXEAGjKS2UOKqI=
github.com/aws/aws-sdk-go-v2/config v1.32.16 h1:Q0iQ7quUgJP0F/SCRTieScnaMdXr9h/2+wze1u3cNeM=
Expand All @@ -74,12 +74,14 @@ github.com/aws/aws-sdk-go-v2/credentials v1.19.15 h1:fyvgWTszojq8hEnMi8PPBTvZdTt
github.com/aws/aws-sdk-go-v2/credentials v1.19.15/go.mod h1:gJiYyMOjNg8OEdRWOf3CrFQxM2a98qmrtjx1zuiQfB8=
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.22 h1:IOGsJ1xVWhsi+ZO7/NW8OuZZBtMJLZbk4P5HDjJO0jQ=
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.22/go.mod h1:b+hYdbU+jGKfXE8kKM6g1+h+L/Go3vMvzlxBsiuGsxg=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.22 h1:GmLa5Kw1ESqtFpXsx5MmC84QWa/ZrLZvlJGa2y+4kcQ=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.22/go.mod h1:6sW9iWm9DK9YRpRGga/qzrzNLgKpT2cIxb7Vo2eNOp0=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.22 h1:dY4kWZiSaXIzxnKlj17nHnBcXXBfac6UlsAx2qL6XrU=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.22/go.mod h1:KIpEUx0JuRZLO7U6cbV204cWAEco2iC3l061IxlwLtI=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.30 h1:xM/Is9cKMHa8Jj8zkvWhvrFkZsXJV9E+BB4g0HW0duQ=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.30/go.mod h1:WueJeNDZvK1fMYEWJIkcivBfEzUkTpBhzlrUKKY8EuA=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.30 h1:jn46zC9LdsVR/ZpMIJqMqb8hHv31BlLx3ulVqNspUOk=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.30/go.mod h1:1hTMsAgbdS/AtUi4bw8+gUuh1pceo+eXRLfpSuSQj3M=
github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.23 h1:FPXsW9+gMuIeKmz7j6ENWcWtBGTe1kH8r9thNt5Uxx4=
github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.23/go.mod h1:7J8iGMdRKk6lw2C+cMIphgAnT8uTwBwNOsGkyOCm80U=
github.com/aws/aws-sdk-go-v2/service/bedrock v1.64.2 h1:TYXv5HG0jLpgeLa8huezwTCTJS2VKnQknkVyRIEZZSY=
github.com/aws/aws-sdk-go-v2/service/bedrock v1.64.2/go.mod h1:pYNYOEFQKBsKwkNQZjVwEuPFTkmSLvAsbSd0HZUwiDw=
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.8 h1:HtOTYcbVcGABLOVuPYaIihj6IlkqubBwFj10K5fxRek=
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.8/go.mod h1:VsK9abqQeGlzPgUr+isNWzPlK2vKe9INMLWnY65f5Xs=
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.22 h1:PUmZeJU6Y1Lbvt9WFuJ0ugUK2xn6hIWUBBbKuOWF30s=
Expand All @@ -94,8 +96,8 @@ github.com/aws/aws-sdk-go-v2/service/ssooidc v1.35.20 h1:oK/njaL8GtyEihkWMD4k3Vg
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.35.20/go.mod h1:JHs8/y1f3zY7U5WcuzoJ/yAYGYtNIVPKLIbp61euvmg=
github.com/aws/aws-sdk-go-v2/service/sts v1.42.0 h1:ks8KBcZPh3PYISr5dAiXCM5/Thcuxk8l+PG4+A0exds=
github.com/aws/aws-sdk-go-v2/service/sts v1.42.0/go.mod h1:pFw33T0WLvXU3rw1WBkpMlkgIn54eCB5FYLhjDc9Foo=
github.com/aws/smithy-go v1.25.0 h1:Sz/XJ64rwuiKtB6j98nDIPyYrV1nVNJ4YU74gttcl5U=
github.com/aws/smithy-go v1.25.0/go.mod h1:YE2RhdIuDbA5E5bTdciG9KrW3+TiEONeUWCqxX9i1Fc=
github.com/aws/smithy-go v1.27.3 h1:F3Zb497UhhskkfpJmfkXswyo+t0sh9OTBnIHjogWbVY=
github.com/aws/smithy-go v1.27.3/go.mod h1:YE2RhdIuDbA5E5bTdciG9KrW3+TiEONeUWCqxX9i1Fc=
github.com/bahlo/generic-list-go v0.2.0 h1:5sz/EEAK+ls5wF+NeqDpk5+iNdMDXrh3z3nPnH1Wvgk=
github.com/bahlo/generic-list-go v0.2.0/go.mod h1:2KvAjgMlE5NNynlg/5iLrrCCZ2+5xWbdbCW3pNTGyYg=
github.com/beorn7/perks v0.0.0-20180321164747-3a771d992973/go.mod h1:Dwedo/Wpr24TaqPxmxbtue+5NUziq4I4S80YR8gNf3Q=
Expand Down
116 changes: 116 additions & 0 deletions provider/bedrock_models_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,116 @@
package provider

import (
"context"
"errors"
"testing"

awsbedrock "github.com/aws/aws-sdk-go-v2/service/bedrock"
"github.com/aws/aws-sdk-go-v2/service/bedrock/types"
)

type fakeBedrockModelLister struct {
input *awsbedrock.ListFoundationModelsInput
out *awsbedrock.ListFoundationModelsOutput
err error
}

func (f *fakeBedrockModelLister) ListFoundationModels(ctx context.Context, in *awsbedrock.ListFoundationModelsInput, optFns ...func(*awsbedrock.Options)) (*awsbedrock.ListFoundationModelsOutput, error) {
f.input = in
return f.out, f.err
}

func TestListBedrockModelsFromAPIListsAnthropicTextModels(t *testing.T) {
streaming := true
api := &fakeBedrockModelLister{out: &awsbedrock.ListFoundationModelsOutput{
ModelSummaries: []types.FoundationModelSummary{
{
ModelId: strPtr("anthropic.claude-sonnet-4-20250514-v1:0"),
ModelName: strPtr("Claude Sonnet 4"),
ProviderName: strPtr("Anthropic"),
OutputModalities: []types.ModelModality{types.ModelModalityText},
InferenceTypesSupported: []types.InferenceType{types.InferenceTypeOnDemand},
ResponseStreamingSupported: &streaming,
},
{
ModelId: strPtr("amazon.titan-text-lite-v1"),
ModelName: strPtr("Titan Text Lite"),
ProviderName: strPtr("Amazon"),
OutputModalities: []types.ModelModality{types.ModelModalityText},
InferenceTypesSupported: []types.InferenceType{types.InferenceTypeOnDemand},
ResponseStreamingSupported: &streaming,
},
{
ModelId: strPtr("anthropic.claude-image-only"),
ModelName: strPtr("Claude Image Only"),
ProviderName: strPtr("Anthropic"),
OutputModalities: []types.ModelModality{types.ModelModalityImage},
InferenceTypesSupported: []types.InferenceType{types.InferenceTypeOnDemand},
},
},
}}

models, err := listBedrockModelsFromAPI(context.Background(), api)
if err != nil {
t.Fatalf("listBedrockModelsFromAPI: %v", err)
}
if api.input == nil {
t.Fatal("ListFoundationModels was not called")
}
if api.input.ByProvider == nil || *api.input.ByProvider != "Anthropic" {
t.Fatalf("ByProvider = %#v", api.input.ByProvider)
}
if api.input.ByOutputModality != types.ModelModalityText {
t.Fatalf("ByOutputModality = %q", api.input.ByOutputModality)
}
if api.input.ByInferenceType != types.InferenceTypeOnDemand {
t.Fatalf("ByInferenceType = %q", api.input.ByInferenceType)
}
if len(models) != 1 {
t.Fatalf("models = %+v", models)
}
if models[0].ID != "anthropic.claude-sonnet-4-20250514-v1:0" {
t.Fatalf("model ID = %q", models[0].ID)
}
if models[0].Name != "Anthropic Claude Sonnet 4" {
t.Fatalf("model name = %q", models[0].Name)
}
}

func TestListBedrockModelsFromAPIReturnsErrors(t *testing.T) {
api := &fakeBedrockModelLister{err: errors.New("access denied")}
_, err := listBedrockModelsFromAPI(context.Background(), api)
if err == nil {
t.Fatal("expected error")
}
}

func TestListModelsWithSettingsBedrockRequiresCredentials(t *testing.T) {
_, err := ListModelsWithSettings(context.Background(), "anthropic_bedrock", "", "", nil)
if err == nil {
t.Fatal("expected missing credentials error")
}
}

func TestBedrockModelListConfigFromJSONCredentials(t *testing.T) {
cfg, err := bedrockModelListConfigFromRequest(ModelListRequest{
APIKey: `{"region":"us-west-2","access_key_id":"AKIAEXAMPLE","secret_access_key":"secret","session_token":"token"}`,
})
if err != nil {
t.Fatalf("bedrockModelListConfigFromRequest: %v", err)
}
if cfg.Region != "us-west-2" {
t.Fatalf("Region = %q", cfg.Region)
}
if cfg.AccessKeyID != "AKIAEXAMPLE" {
t.Fatalf("AccessKeyID = %q", cfg.AccessKeyID)
}
if cfg.SecretAccessKey != "secret" {
t.Fatalf("SecretAccessKey = %q", cfg.SecretAccessKey)
}
if cfg.SessionToken != "token" {
t.Fatalf("SessionToken = %q", cfg.SessionToken)
}
}

func strPtr(s string) *string { return &s }
155 changes: 152 additions & 3 deletions provider/models.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,10 @@ import (
"sort"
"strings"

"github.com/aws/aws-sdk-go-v2/aws"
"github.com/aws/aws-sdk-go-v2/credentials"
awsbedrock "github.com/aws/aws-sdk-go-v2/service/bedrock"
bedrocktypes "github.com/aws/aws-sdk-go-v2/service/bedrock/types"
"github.com/google/generative-ai-go/genai"
"google.golang.org/api/iterator"
googleoption "google.golang.org/api/option"
Expand All @@ -26,6 +30,7 @@ type ModelListRequest struct {
ProviderType string
APIKey string
BaseURL string
Settings map[string]string
}

type ModelLister func(ctx context.Context, req ModelListRequest) ([]ModelInfo, error)
Expand Down Expand Up @@ -53,11 +58,15 @@ type copilotTokenResponse struct {
// ListModels fetches available models from the given provider type.
// Only requires an API key and optional base URL — no saved provider needed.
func ListModels(ctx context.Context, providerType, apiKey, baseURL string) ([]ModelInfo, error) {
return ListModelsWithSettings(ctx, providerType, apiKey, baseURL, nil)
}

func ListModelsWithSettings(ctx context.Context, providerType, apiKey, baseURL string, settings map[string]string) ([]ModelInfo, error) {
lister, ok := modelListers[providerType]
if !ok {
return nil, fmt.Errorf("unsupported provider type: %s", providerType)
}
return lister(ctx, ModelListRequest{ProviderType: providerType, APIKey: apiKey, BaseURL: baseURL})
return lister(ctx, ModelListRequest{ProviderType: providerType, APIKey: apiKey, BaseURL: baseURL, Settings: settings})
}

var modelListers = map[string]ModelLister{
Expand Down Expand Up @@ -90,8 +99,8 @@ var modelListers = map[string]ModelLister{
"openai_azure": func(context.Context, ModelListRequest) ([]ModelInfo, error) {
return nil, dynamicModelListingUnsupported("openai_azure")
},
"anthropic_bedrock": func(context.Context, ModelListRequest) ([]ModelInfo, error) {
return nil, dynamicModelListingUnsupported("anthropic_bedrock")
"anthropic_bedrock": func(ctx context.Context, req ModelListRequest) ([]ModelInfo, error) {
return listBedrockModels(ctx, req)
},
"anthropic_vertex": func(context.Context, ModelListRequest) ([]ModelInfo, error) {
return nil, dynamicModelListingUnsupported("anthropic_vertex")
Expand Down Expand Up @@ -604,6 +613,146 @@ func dynamicModelListingUnsupported(providerType string) error {
return fmt.Errorf("%s: dynamic model listing is not implemented for this provider; specify a custom model ID", providerType)
}

type bedrockModelListConfig struct {
Region string
AccessKeyID string
SecretAccessKey string
SessionToken string
}

type bedrockFoundationModelLister interface {
ListFoundationModels(context.Context, *awsbedrock.ListFoundationModelsInput, ...func(*awsbedrock.Options)) (*awsbedrock.ListFoundationModelsOutput, error)
}

func bedrockModelListConfigFromRequest(req ModelListRequest) (bedrockModelListConfig, error) {
cfg := bedrockModelListConfig{
Region: req.Settings["region"],
AccessKeyID: req.Settings["access_key_id"],
SecretAccessKey: req.APIKey,
SessionToken: req.Settings["session_token"],
}
if strings.TrimSpace(req.APIKey) != "" && strings.HasPrefix(strings.TrimSpace(req.APIKey), "{") {
var secret struct {
Region string `json:"region"`
AccessKeyID string `json:"access_key_id"`
SecretAccessKey string `json:"secret_access_key"`
SessionToken string `json:"session_token"`
}
if err := json.Unmarshal([]byte(req.APIKey), &secret); err != nil {
return bedrockModelListConfig{}, fmt.Errorf("anthropic_bedrock: parse credential JSON: %w", err)
}
cfg.SecretAccessKey = ""
if cfg.Region == "" {
cfg.Region = secret.Region
}
if cfg.AccessKeyID == "" {
cfg.AccessKeyID = secret.AccessKeyID
}
if cfg.SecretAccessKey == "" {
cfg.SecretAccessKey = secret.SecretAccessKey
}
if cfg.SessionToken == "" {
cfg.SessionToken = secret.SessionToken
}
}
if cfg.Region == "" {
cfg.Region = defaultBedrockRegion
}
cfg.AccessKeyID = strings.TrimSpace(cfg.AccessKeyID)
cfg.SecretAccessKey = strings.TrimSpace(cfg.SecretAccessKey)
cfg.SessionToken = strings.TrimSpace(cfg.SessionToken)
if cfg.AccessKeyID == "" {
return bedrockModelListConfig{}, fmt.Errorf("anthropic_bedrock: access_key_id setting is required for model discovery")
}
if cfg.SecretAccessKey == "" {
return bedrockModelListConfig{}, fmt.Errorf("anthropic_bedrock: secret access key is required for model discovery")
}
return cfg, nil
}

func listBedrockModels(ctx context.Context, req ModelListRequest) ([]ModelInfo, error) {
if err := ValidateBaseURL(req.BaseURL); err != nil {
return nil, err
}
cfg, err := bedrockModelListConfigFromRequest(req)
if err != nil {
return nil, err
}
awsCfg := aws.Config{
Region: cfg.Region,
Credentials: credentials.NewStaticCredentialsProvider(cfg.AccessKeyID, cfg.SecretAccessKey, cfg.SessionToken),
HTTPClient: modelHTTPClient,
}
opts := []func(*awsbedrock.Options){}
if req.BaseURL != "" {
opts = append(opts, func(o *awsbedrock.Options) {
o.BaseEndpoint = aws.String(strings.TrimRight(req.BaseURL, "/"))
})
}
return listBedrockModelsFromAPI(ctx, awsbedrock.NewFromConfig(awsCfg, opts...))
}

func listBedrockModelsFromAPI(ctx context.Context, api bedrockFoundationModelLister) ([]ModelInfo, error) {
out, err := api.ListFoundationModels(ctx, &awsbedrock.ListFoundationModelsInput{
ByProvider: aws.String("Anthropic"),
ByOutputModality: bedrocktypes.ModelModalityText,
ByInferenceType: bedrocktypes.InferenceTypeOnDemand,
})
if err != nil {
return nil, fmt.Errorf("anthropic_bedrock: list foundation models: %w", err)
}
models := make([]ModelInfo, 0, len(out.ModelSummaries))
for _, summary := range out.ModelSummaries {
id := strings.TrimSpace(aws.ToString(summary.ModelId))
if id == "" {
continue
}
if providerName := strings.TrimSpace(aws.ToString(summary.ProviderName)); providerName != "" && !strings.EqualFold(providerName, "Anthropic") {
continue
}
if len(summary.OutputModalities) > 0 && !modelModalitiesContain(summary.OutputModalities, bedrocktypes.ModelModalityText) {
continue
}
if len(summary.InferenceTypesSupported) > 0 && !inferenceTypesContain(summary.InferenceTypesSupported, bedrocktypes.InferenceTypeOnDemand) {
continue
}
name := strings.TrimSpace(aws.ToString(summary.ModelName))
if name == "" {
name = id
}
providerName := strings.TrimSpace(aws.ToString(summary.ProviderName))
if providerName != "" && !strings.HasPrefix(strings.ToLower(name), strings.ToLower(providerName)+" ") {
name = providerName + " " + name
}
models = append(models, ModelInfo{ID: id, Name: name})
}
if len(models) == 0 {
return nil, fmt.Errorf("anthropic_bedrock: no selectable Anthropic text models returned")
}
sort.Slice(models, func(i, j int) bool {
return models[i].ID < models[j].ID
})
return models, nil
}

func modelModalitiesContain(values []bedrocktypes.ModelModality, want bedrocktypes.ModelModality) bool {
for _, v := range values {
if v == want {
return true
}
}
return false
}

func inferenceTypesContain(values []bedrocktypes.InferenceType, want bedrocktypes.InferenceType) bool {
for _, v := range values {
if v == want {
return true
}
}
return false
}

// listOllamaModels lists models available on a local Ollama server.
func listOllamaModels(ctx context.Context, baseURL string) ([]ModelInfo, error) {
c := NewOllamaClient(baseURL)
Expand Down
2 changes: 1 addition & 1 deletion provider/models_dynamic_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@ import (
)

func TestListModelsHostedProviderVariantsRequireDynamicCatalog(t *testing.T) {
for _, providerType := range []string{"openai_azure", "anthropic_bedrock", "anthropic_vertex", "anthropic_foundry"} {
for _, providerType := range []string{"openai_azure", "anthropic_vertex", "anthropic_foundry"} {
t.Run(providerType, func(t *testing.T) {
_, err := ListModels(context.Background(), providerType, "", "")
if err == nil {
Expand Down
Loading