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
112 changes: 101 additions & 11 deletions cmd/ratchet/cmd_provider.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package main
import (
"bufio"
"context"
"encoding/json"
"fmt"
"io"
"os"
Expand Down Expand Up @@ -72,7 +73,8 @@ func handleProvider(args []string) {
}
// Parse --model flag from remaining args.
model, modelSet := parseProviderModelFlag(args)
var apiKey, baseURL string
var apiKey, baseURL, settingsJSON string
var settings map[string]string
scanner := bufio.NewScanner(os.Stdin)
switch providerType {
case "ollama":
Expand Down Expand Up @@ -121,6 +123,17 @@ func handleProvider(args []string) {
case "openai_chatgpt":
fmt.Println("Use: ratchet provider setup openai-chatgpt")
return
case "anthropic_bedrock":
apiKey, settings, err = promptBedrockProviderCredentials(scanner, os.Stdout, providerauth.PromptSecret)
if err != nil {
fmt.Fprintf(os.Stderr, "error: %v\n", err)
os.Exit(1)
}
settingsJSON, err = providerSettingsJSON(settings)
if err != nil {
fmt.Fprintf(os.Stderr, "error: %v\n", err)
os.Exit(1)
}
default:
apiKey, err = providerauth.PromptAPIKey(providerType)
if err != nil {
Expand All @@ -137,7 +150,7 @@ func handleProvider(args []string) {
}
if !modelSet && model == "" {
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
selected, selectErr := promptProviderModelSelection(ctx, providerType, apiKey, baseURL, scanner, os.Stdout, wfprovider.ListModels)
selected, selectErr := promptProviderModelSelection(ctx, providerType, apiKey, baseURL, settings, scanner, os.Stdout, wfprovider.ListModelsWithSettings)
cancel()
if selectErr != nil {
fmt.Fprintf(os.Stderr, "model selection failed: %v\n", selectErr)
Expand All @@ -146,11 +159,12 @@ func handleProvider(args []string) {
model = selected
}
p, err := c.AddProvider(context.Background(), &pb.AddProviderReq{
Alias: alias,
Type: providerType,
Model: model,
ApiKey: apiKey,
BaseUrl: baseURL,
Alias: alias,
Type: providerType,
Model: model,
ApiKey: apiKey,
BaseUrl: baseURL,
Settings: settingsJSON,
})
if err != nil {
fmt.Fprintf(os.Stderr, "error: %v\n", err)
Expand Down Expand Up @@ -298,7 +312,7 @@ func handleOpenAIChatGPTSetup(args []string) {
scanner := bufio.NewScanner(os.Stdin)
if !opts.modelSet {
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
model, selectErr := promptProviderModelSelection(ctx, "openai_chatgpt", tokenBundle, "", scanner, os.Stdout, wfprovider.ListModels)
model, selectErr := promptProviderModelSelection(ctx, "openai_chatgpt", tokenBundle, "", nil, scanner, os.Stdout, wfprovider.ListModelsWithSettings)
cancel()
if selectErr != nil {
fmt.Fprintf(os.Stderr, "model selection failed: %v\n", selectErr)
Expand Down Expand Up @@ -689,10 +703,86 @@ func promptModelSelection(models []wfprovider.ModelInfo, scanner *bufio.Scanner)
return models[0].ID
}

type providerModelLister func(context.Context, string, string, string) ([]wfprovider.ModelInfo, error)
type providerModelLister func(context.Context, string, string, string, map[string]string) ([]wfprovider.ModelInfo, error)

type providerSecretPrompter func(string) (string, error)

func promptBedrockProviderCredentials(scanner *bufio.Scanner, out io.Writer, promptSecret providerSecretPrompter) (string, map[string]string, error) {
fmt.Fprint(out, "AWS access key ID: ")
if !scanner.Scan() {
if scanErr := scanner.Err(); scanErr != nil {
return "", nil, scanErr
}
return "", nil, io.EOF
}
accessKeyID := scanner.Text()

secretAccessKey, err := promptSecret("AWS secret access key")
if err != nil {
return "", nil, err
}
secretAccessKey = strings.TrimSpace(secretAccessKey)
if secretAccessKey == "" {
return "", nil, fmt.Errorf("AWS secret access key is required")
}

fmt.Fprint(out, "AWS region [us-east-1]: ")
if !scanner.Scan() {
if scanErr := scanner.Err(); scanErr != nil {
return "", nil, scanErr
}
return "", nil, io.EOF
}
region := scanner.Text()

settings, err := bedrockProviderSettings(accessKeyID, region)
if err != nil {
return "", nil, err
}
return secretAccessKey, settings, nil
}

func bedrockProviderSettings(accessKeyID, region string) (map[string]string, error) {
accessKeyID = strings.TrimSpace(accessKeyID)
region = strings.TrimSpace(region)
if accessKeyID == "" {
return nil, fmt.Errorf("AWS access key ID is required")
}
if region == "" {
region = "us-east-1"
}
settings := map[string]string{
"access_key_id": accessKeyID,
"region": region,
}
return settings, nil
}

func providerSettingsJSON(settings map[string]string) (string, error) {
if len(settings) == 0 {
return "", nil
}
clean := make(map[string]string, len(settings))
for k, v := range settings {
k = strings.TrimSpace(k)
v = strings.TrimSpace(v)
if k == "" || v == "" {
continue
}
clean[k] = v
}
if len(clean) == 0 {
return "", nil
}
data, err := json.Marshal(clean)
if err != nil {
return "", fmt.Errorf("marshal provider settings: %w", err)
}
return string(data), nil
}

func promptProviderModelSelection(ctx context.Context, providerType, apiKey, baseURL string, scanner *bufio.Scanner, out io.Writer, list providerModelLister) (string, error) {
models, err := list(ctx, providerType, apiKey, baseURL)
func promptProviderModelSelection(ctx context.Context, providerType, apiKey, baseURL string, settings map[string]string, scanner *bufio.Scanner, out io.Writer, list providerModelLister) (string, error) {
models, err := list(ctx, providerType, apiKey, baseURL, settings)
if err != nil || len(models) == 0 {
if err != nil {
fmt.Fprintf(out, "could not list models for %s: %v\n", providerType, err)
Expand Down
88 changes: 85 additions & 3 deletions cmd/ratchet/cmd_provider_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package main
import (
"bufio"
"context"
"encoding/json"
"errors"
"os"
"runtime"
Expand Down Expand Up @@ -150,16 +151,95 @@ func TestOpenAIChatGPTAddProviderReq(t *testing.T) {
}
}

func TestProviderSettingsJSON(t *testing.T) {
settingsJSON, err := providerSettingsJSON(map[string]string{
"region": "us-west-2",
"access_key_id": "AKIAEXAMPLE",
"session_token": "",
})
if err != nil {
t.Fatalf("providerSettingsJSON: %v", err)
}
var settings map[string]string
if err := json.Unmarshal([]byte(settingsJSON), &settings); err != nil {
t.Fatalf("settings JSON is invalid: %v", err)
}
if settings["region"] != "us-west-2" {
t.Fatalf("region = %q", settings["region"])
}
if settings["access_key_id"] != "AKIAEXAMPLE" {
t.Fatalf("access_key_id = %q", settings["access_key_id"])
}
if _, ok := settings["session_token"]; ok {
t.Fatalf("empty session_token should be omitted: %#v", settings)
}
}

func TestBedrockProviderSettingsDefaultsRegion(t *testing.T) {
settings, err := bedrockProviderSettings(" AKIAEXAMPLE ", "")
if err != nil {
t.Fatalf("bedrockProviderSettings: %v", err)
}
if settings["access_key_id"] != "AKIAEXAMPLE" {
t.Fatalf("access_key_id = %q", settings["access_key_id"])
}
if settings["region"] != "us-east-1" {
t.Fatalf("region = %q", settings["region"])
}
if _, ok := settings["session_token"]; ok {
t.Fatalf("session_token should not be stored in settings: %#v", settings)
}
}

func TestPromptBedrockProviderCredentials(t *testing.T) {
scanner := bufio.NewScanner(strings.NewReader("AKIAEXAMPLE\nus-west-2\n"))
apiKey, settings, err := promptBedrockProviderCredentials(scanner, &strings.Builder{}, func(label string) (string, error) {
if label != "AWS secret access key" {
t.Fatalf("label = %q", label)
}
return " secret ", nil
})
if err != nil {
t.Fatalf("promptBedrockProviderCredentials: %v", err)
}
if apiKey != "secret" {
t.Fatalf("apiKey = %q", apiKey)
}
if settings["access_key_id"] != "AKIAEXAMPLE" || settings["region"] != "us-west-2" {
t.Fatalf("settings = %#v", settings)
}
if _, ok := settings["session_token"]; ok {
t.Fatalf("session_token should not be stored in settings: %#v", settings)
}
}

func TestPromptBedrockProviderCredentialsRequiresSecret(t *testing.T) {
scanner := bufio.NewScanner(strings.NewReader("AKIAEXAMPLE\nus-west-2\n"))
_, _, err := promptBedrockProviderCredentials(scanner, &strings.Builder{}, func(string) (string, error) {
return " ", nil
})
if err == nil {
t.Fatal("expected missing secret error")
}
}

func TestPromptProviderModelSelectionDefaultsToFirstEnumeratedModel(t *testing.T) {
scanner := bufio.NewScanner(strings.NewReader("\n"))
model, err := promptProviderModelSelection(
context.Background(),
"openai",
"api-key",
"",
map[string]string{"region": "us-east-1"},
scanner,
&strings.Builder{},
func(context.Context, string, string, string) ([]wfprovider.ModelInfo, error) {
func(_ context.Context, providerType, apiKey, baseURL string, settings map[string]string) ([]wfprovider.ModelInfo, error) {
if providerType != "openai" || apiKey != "api-key" || baseURL != "" {
t.Fatalf("unexpected lister args: %q %q %q", providerType, apiKey, baseURL)
}
if settings["region"] != "us-east-1" {
t.Fatalf("settings not passed to lister: %#v", settings)
}
return []wfprovider.ModelInfo{
{ID: "gpt-5.5", Name: "GPT-5.5"},
{ID: "gpt-5.4-mini", Name: "GPT-5.4-Mini"},
Expand All @@ -181,9 +261,10 @@ func TestPromptProviderModelSelectionSupportsCustomModel(t *testing.T) {
"openai",
"api-key",
"",
nil,
scanner,
&strings.Builder{},
func(context.Context, string, string, string) ([]wfprovider.ModelInfo, error) {
func(context.Context, string, string, string, map[string]string) ([]wfprovider.ModelInfo, error) {
return []wfprovider.ModelInfo{
{ID: "gpt-5.5", Name: "GPT-5.5"},
{ID: "gpt-5.4-mini", Name: "GPT-5.4-Mini"},
Expand All @@ -206,9 +287,10 @@ func TestPromptProviderModelSelectionPromptsManualWhenEnumerationFails(t *testin
"anthropic_bedrock",
"",
"",
nil,
scanner,
&out,
func(context.Context, string, string, string) ([]wfprovider.ModelInfo, error) {
func(context.Context, string, string, string, map[string]string) ([]wfprovider.ModelInfo, error) {
return nil, errors.New("no dynamic catalog")
},
)
Expand Down
11 changes: 6 additions & 5 deletions go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ require (
github.com/GoCodeAlone/acpx-go v0.2.1
github.com/GoCodeAlone/workflow v0.85.2
github.com/GoCodeAlone/workflow-plugin-acpx v0.1.0
github.com/GoCodeAlone/workflow-plugin-agent v0.12.4
github.com/GoCodeAlone/workflow-plugin-agent v0.12.5
github.com/charmbracelet/glamour v0.10.0
github.com/coder/acp-go-sdk v0.6.3
github.com/creack/pty v1.1.24
Expand Down Expand Up @@ -53,22 +53,23 @@ require (
github.com/anthropics/anthropic-sdk-go v1.26.0 // indirect
github.com/armon/go-metrics v0.4.1 // indirect
github.com/atotto/clipboard v0.1.4 // indirect
github.com/aws/aws-sdk-go-v2 v1.41.6 // indirect
github.com/aws/aws-sdk-go-v2 v1.42.1 // indirect
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/credentials v1.19.15 // 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/bedrock v1.64.2 // 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
github.com/aws/aws-sdk-go-v2/service/kinesis v1.43.5 // indirect
github.com/aws/aws-sdk-go-v2/service/signin v1.0.10 // indirect
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/aymanbagabas/go-osc52/v2 v2.0.1 // indirect
github.com/aymerick/douceur v0.2.0 // indirect
github.com/bahlo/generic-list-go v0.2.0 // indirect
Expand Down
22 changes: 12 additions & 10 deletions go.sum
Original file line number Diff line number Diff line change
Expand Up @@ -56,8 +56,8 @@ github.com/GoCodeAlone/workflow v0.85.2 h1:u65GfzC0c1bkicOOvyPbgcB1rhodexWd7f+wR
github.com/GoCodeAlone/workflow v0.85.2/go.mod h1:tqWHdWHCn2ESB4Rj6FiRjKchof9zjcs11o8AtfOiZaU=
github.com/GoCodeAlone/workflow-plugin-acpx v0.1.0 h1:mwwaLAfsxEIYdIQ4wvV/PkT0Ol30ntGhiun6vArRexI=
github.com/GoCodeAlone/workflow-plugin-acpx v0.1.0/go.mod h1:Bun12gpt+/WmqlAafEXpxSG1esbI3TAmwWFrgaR1xOY=
github.com/GoCodeAlone/workflow-plugin-agent v0.12.4 h1:hwql/brpEjIlL/arQ+eV/OP5vuxFtMdo7TV/+X14Ogg=
github.com/GoCodeAlone/workflow-plugin-agent v0.12.4/go.mod h1:YKSaEo28dYbP2UkrRVjEcrIToMAXfLkTrtlz6Sbv4m8=
github.com/GoCodeAlone/workflow-plugin-agent v0.12.5 h1:RwduRH/Fe7BKs91z9bw2VxhRo399pe/BGt5kcdzmUyY=
github.com/GoCodeAlone/workflow-plugin-agent v0.12.5/go.mod h1:lYhGZ6Xh1mLtaJq68aWzpcNA2wH+L/q7cP46lujUre8=
github.com/GoCodeAlone/workflow-plugin-authz v0.5.13 h1:5vORZ5sUsVWEs4r/zmmy1cYSFGh8ZzhbVbgn0UhhoXs=
github.com/GoCodeAlone/workflow-plugin-authz v0.5.13/go.mod h1:O5XEgNZmkqDN+bFGXWI/DUsav0v8brELSbEdxlW9lPQ=
github.com/GoCodeAlone/yaegi v0.17.2 h1:WK6Y6e0t1a6U7r+S2dN3CGWW1PizYD3zO0zneToZPxM=
Expand Down Expand Up @@ -101,8 +101,8 @@ github.com/atotto/clipboard v0.1.4 h1:EH0zSVneZPSuFR11BlR9YppQTVDbh5+16AmcJi4g1z
github.com/atotto/clipboard v0.1.4/go.mod h1:ZY9tmq7sm5xIbd9bOK4onWV4S6X0u6GY7Vn0Yu86PYI=
github.com/autarch/testify v1.2.2 h1:9Q9V6zqhP7R6dv+zRUddv6kXKLo6ecQhnFRFWM71i1c=
github.com/autarch/testify v1.2.2/go.mod h1:oDbHKfFv2/D5UtVrxkk90OKcb6P4/AqF1Pcf6ZbvDQo=
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 @@ -111,12 +111,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 @@ -131,8 +133,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/aymanbagabas/go-osc52/v2 v2.0.1 h1:HwpRHbFMcZLEVr42D4p7XBqjyuxQH5SMiErDT4WkJ2k=
github.com/aymanbagabas/go-osc52/v2 v2.0.1/go.mod h1:uYgXzlJ7ZpABp8OJ+exZzJJhRNQ2ASbcXHWsFqH8hp8=
github.com/aymanbagabas/go-udiff v0.4.1 h1:OEIrQ8maEeDBXQDoGCbbTTXYJMYRCRO1fnodZ12Gv5o=
Expand Down
Loading
Loading