diff --git a/cmd/ratchet/cmd_provider.go b/cmd/ratchet/cmd_provider.go index 8decc3b..852cf73 100644 --- a/cmd/ratchet/cmd_provider.go +++ b/cmd/ratchet/cmd_provider.go @@ -3,6 +3,7 @@ package main import ( "bufio" "context" + "encoding/json" "fmt" "io" "os" @@ -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": @@ -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 { @@ -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) @@ -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) @@ -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) @@ -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) diff --git a/cmd/ratchet/cmd_provider_test.go b/cmd/ratchet/cmd_provider_test.go index 959602d..dbbda35 100644 --- a/cmd/ratchet/cmd_provider_test.go +++ b/cmd/ratchet/cmd_provider_test.go @@ -3,6 +3,7 @@ package main import ( "bufio" "context" + "encoding/json" "errors" "os" "runtime" @@ -150,6 +151,78 @@ 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( @@ -157,9 +230,16 @@ func TestPromptProviderModelSelectionDefaultsToFirstEnumeratedModel(t *testing.T "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"}, @@ -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"}, @@ -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") }, ) diff --git a/go.mod b/go.mod index 3ec7b47..318ad3f 100644 --- a/go.mod +++ b/go.mod @@ -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 @@ -53,14 +53,15 @@ 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 @@ -68,7 +69,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/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 diff --git a/go.sum b/go.sum index 9cb055e..5fd1142 100644 --- a/go.sum +++ b/go.sum @@ -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= @@ -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= @@ -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= @@ -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= diff --git a/internal/daemon/e2e_provider_test.go b/internal/daemon/e2e_provider_test.go index 20a7032..15da941 100644 --- a/internal/daemon/e2e_provider_test.go +++ b/internal/daemon/e2e_provider_test.go @@ -195,6 +195,66 @@ func TestE2EAddProviderWithoutKey_NoSecret(t *testing.T) { } } +func TestE2EAddProviderStoresSettings(t *testing.T) { + h := newE2EHarness(t) + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + + _, err := h.Client.AddProvider(ctx, &pb.AddProviderReq{ + Alias: "bedrock", + Type: "anthropic_bedrock", + Model: "anthropic.claude-sonnet-4-20250514-v1:0", + ApiKey: "secret", + Settings: `{"region":"us-west-2","access_key_id":"AKIAEXAMPLE"}`, + }) + if err != nil { + t.Fatalf("AddProvider: %v", err) + } + + var settings string + if err := h.DB.QueryRowContext(ctx, + `SELECT settings FROM llm_providers WHERE alias = ?`, "bedrock", + ).Scan(&settings); err != nil { + t.Fatalf("query settings: %v", err) + } + if settings != `{"region":"us-west-2","access_key_id":"AKIAEXAMPLE"}` { + t.Fatalf("settings = %q", settings) + } + + _, err = h.Client.AddProvider(ctx, &pb.AddProviderReq{ + Alias: "bedrock", + Type: "anthropic_bedrock", + Model: "anthropic.claude-opus-4-20250514-v1:0", + Settings: `{ }`, + }) + if err != nil { + t.Fatalf("AddProvider update: %v", err) + } + if err := h.DB.QueryRowContext(ctx, + `SELECT settings FROM llm_providers WHERE alias = ?`, "bedrock", + ).Scan(&settings); err != nil { + t.Fatalf("query updated settings: %v", err) + } + if settings != `{"region":"us-west-2","access_key_id":"AKIAEXAMPLE"}` { + t.Fatalf("settings after update = %q", settings) + } +} + +func TestE2EAddProviderRejectsInvalidSettings(t *testing.T) { + h := newE2EHarness(t) + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + + _, err := h.Client.AddProvider(ctx, &pb.AddProviderReq{ + Alias: "bad-settings", + Type: "anthropic_bedrock", + Settings: `["not","an","object"]`, + }) + if err == nil { + t.Fatal("expected invalid settings error") + } +} + // TestE2EStaleMigration_KeylessTypesGetEmptySecretName verifies the initDB // migration that clears stale secret_name values for ollama/llama_cpp rows that // were inserted before the fix. We inject a stale row directly into the DB diff --git a/internal/daemon/service.go b/internal/daemon/service.go index b8fec8b..0738c32 100644 --- a/internal/daemon/service.go +++ b/internal/daemon/service.go @@ -563,6 +563,11 @@ func (s *Service) AddProvider(ctx context.Context, req *pb.AddProviderReq) (*pb. maxTokens = 4096 } + settings, err := normalizeProviderSettings(req.Settings) + if err != nil { + return nil, status.Errorf(codes.InvalidArgument, "invalid provider settings: %v", err) + } + // Upsert: insert or update if alias already exists. if _, err := tx.ExecContext(ctx, `INSERT INTO llm_providers (id, alias, type, model, secret_name, base_url, max_tokens, settings, is_default) @@ -572,9 +577,10 @@ func (s *Service) AddProvider(ctx context.Context, req *pb.AddProviderReq) (*pb. model = excluded.model, secret_name = CASE WHEN excluded.secret_name = '' THEN secret_name ELSE excluded.secret_name END, base_url = CASE WHEN excluded.base_url = '' THEN base_url ELSE excluded.base_url END, + settings = CASE WHEN excluded.settings = '{}' THEN settings ELSE excluded.settings END, max_tokens = excluded.max_tokens, is_default = excluded.is_default`, - id, req.Alias, req.Type, req.Model, secretName, req.BaseUrl, maxTokens, "{}", isDefault, + id, req.Alias, req.Type, req.Model, secretName, req.BaseUrl, maxTokens, settings, isDefault, ); err != nil { return nil, status.Errorf(codes.Internal, "insert provider: %v", err) } @@ -604,6 +610,25 @@ func (s *Service) AddProvider(ctx context.Context, req *pb.AddProviderReq) (*pb. }, nil } +func normalizeProviderSettings(settings string) (string, error) { + settings = strings.TrimSpace(settings) + if settings == "" { + return "{}", nil + } + var decoded any + if err := json.Unmarshal([]byte(settings), &decoded); err != nil { + return "", err + } + settingsObject, ok := decoded.(map[string]any) + if !ok { + return "", fmt.Errorf("must be a JSON object") + } + if len(settingsObject) == 0 { + return "{}", nil + } + return settings, nil +} + func validProviderAliasForSecret(alias string) bool { if alias == "" || len(alias) > 128 { return false diff --git a/internal/proto/ratchet.pb.go b/internal/proto/ratchet.pb.go index a484cda..a0b8d9e 100644 --- a/internal/proto/ratchet.pb.go +++ b/internal/proto/ratchet.pb.go @@ -2419,6 +2419,7 @@ type AddProviderReq struct { BaseUrl string `protobuf:"bytes,5,opt,name=base_url,json=baseUrl,proto3" json:"base_url,omitempty"` MaxTokens int32 `protobuf:"varint,6,opt,name=max_tokens,json=maxTokens,proto3" json:"max_tokens,omitempty"` IsDefault bool `protobuf:"varint,7,opt,name=is_default,json=isDefault,proto3" json:"is_default,omitempty"` + Settings string `protobuf:"bytes,8,opt,name=settings,proto3" json:"settings,omitempty"` unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache } @@ -2502,6 +2503,13 @@ func (x *AddProviderReq) GetIsDefault() bool { return false } +func (x *AddProviderReq) GetSettings() string { + if x != nil { + return x.Settings + } + return "" +} + type Provider struct { state protoimpl.MessageState `protogen:"open.v1"` Alias string `protobuf:"bytes,1,opt,name=alias,proto3" json:"alias,omitempty"` @@ -7130,7 +7138,7 @@ const file_internal_proto_ratchet_proto_rawDesc = "" + "\ftool_call_id\x18\x04 \x01(\tR\n" + "toolCallId\x128\n" + "\ttimestamp\x18\x05 \x01(\v2\x1a.google.protobuf.TimestampR\ttimestamp\x12\x0e\n" + - "\x02id\x18\x06 \x01(\tR\x02id\"\xc2\x01\n" + + "\x02id\x18\x06 \x01(\tR\x02id\"\xde\x01\n" + "\x0eAddProviderReq\x12\x14\n" + "\x05alias\x18\x01 \x01(\tR\x05alias\x12\x12\n" + "\x04type\x18\x02 \x01(\tR\x04type\x12\x14\n" + @@ -7140,7 +7148,8 @@ const file_internal_proto_ratchet_proto_rawDesc = "" + "\n" + "max_tokens\x18\x06 \x01(\x05R\tmaxTokens\x12\x1d\n" + "\n" + - "is_default\x18\a \x01(\bR\tisDefault\"\x84\x01\n" + + "is_default\x18\a \x01(\bR\tisDefault\x12\x1a\n" + + "\bsettings\x18\b \x01(\tR\bsettings\"\x84\x01\n" + "\bProvider\x12\x14\n" + "\x05alias\x18\x01 \x01(\tR\x05alias\x12\x12\n" + "\x04type\x18\x02 \x01(\tR\x04type\x12\x14\n" + diff --git a/internal/proto/ratchet.proto b/internal/proto/ratchet.proto index 359ffbe..cb6bace 100644 --- a/internal/proto/ratchet.proto +++ b/internal/proto/ratchet.proto @@ -245,6 +245,7 @@ message AddProviderReq { string base_url = 5; int32 max_tokens = 6; bool is_default = 7; + string settings = 8; } message Provider { diff --git a/internal/provider/auth.go b/internal/provider/auth.go index 037425b..f35dea1 100644 --- a/internal/provider/auth.go +++ b/internal/provider/auth.go @@ -13,7 +13,12 @@ import ( // PromptAPIKey prompts the user for an API key (hidden input). func PromptAPIKey(providerType string) (string, error) { - fmt.Printf("Enter %s API key: ", providerType) + return PromptSecret(fmt.Sprintf("Enter %s API key", providerType)) +} + +// PromptSecret prompts the user for a secret value (hidden input). +func PromptSecret(label string) (string, error) { + fmt.Printf("%s: ", label) key, err := term.ReadPassword(int(syscall.Stdin)) fmt.Println() if err != nil { diff --git a/internal/provider/models.go b/internal/provider/models.go index 43413ab..1bbdc12 100644 --- a/internal/provider/models.go +++ b/internal/provider/models.go @@ -13,3 +13,8 @@ type ModelInfo = wfprovider.ModelInfo func ListModels(ctx context.Context, providerType, apiKey, baseURL string) ([]ModelInfo, error) { return wfprovider.ListModels(ctx, providerType, apiKey, baseURL) } + +// ListModelsWithSettings fetches available models with provider-specific discovery settings. +func ListModelsWithSettings(ctx context.Context, providerType, apiKey, baseURL string, settings map[string]string) ([]ModelInfo, error) { + return wfprovider.ListModelsWithSettings(ctx, providerType, apiKey, baseURL, settings) +} diff --git a/internal/tui/pty_capture_test.go b/internal/tui/pty_capture_test.go index 8aa95ba..3db33f5 100644 --- a/internal/tui/pty_capture_test.go +++ b/internal/tui/pty_capture_test.go @@ -26,7 +26,7 @@ func TestStartTUITestPTYDisablesSoftwareFlowControl(t *testing.T) { } dir := t.TempDir() script := filepath.Join(dir, "show-stty") - if err := os.WriteFile(script, []byte("#!/bin/sh\nstty -a\n"), 0700); err != nil { + if err := os.WriteFile(script, []byte("#!/bin/sh\nsleep 1\nstty -a\n"), 0700); err != nil { t.Fatalf("write stty script: %v", err) }