From 1d1cffcfe464baade262660bdf33406ff1379930 Mon Sep 17 00:00:00 2001 From: Jon Langevin Date: Mon, 6 Jul 2026 09:50:18 -0400 Subject: [PATCH] feat: support generic provider setup --- cmd/ratchet/cmd_provider.go | 65 ++++++++++++++++++++++++++++---- cmd/ratchet/cmd_provider_test.go | 41 ++++++++++++++++++++ go.mod | 5 ++- go.sum | 6 +++ 4 files changed, 108 insertions(+), 9 deletions(-) diff --git a/cmd/ratchet/cmd_provider.go b/cmd/ratchet/cmd_provider.go index 8fc0f0a..ed5d43f 100644 --- a/cmd/ratchet/cmd_provider.go +++ b/cmd/ratchet/cmd_provider.go @@ -209,31 +209,42 @@ func handleProvider(args []string) { case "openai_chatgpt": fmt.Println("Use: ratchet provider setup openai-chatgpt") return - case "anthropic_bedrock": + case "bedrock", "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: + if providerType == "custom" { + settings, err = promptCustomProviderCompatibility(scanner, os.Stdout) + if err != nil { + fmt.Fprintf(os.Stderr, "error: %v\n", err) + os.Exit(1) + } + } apiKey, err = providerauth.PromptAPIKey(providerType) if err != nil { fmt.Fprintf(os.Stderr, "error: %v\n", err) os.Exit(1) } - if providerType == "custom" || providerType == "openai" { + if providerPromptsBaseURL(providerType) { baseURL, err = providerauth.PromptBaseURL("") if err != nil { fmt.Fprintf(os.Stderr, "error: %v\n", err) os.Exit(1) } + if providerRequiresBaseURL(providerType) && strings.TrimSpace(baseURL) == "" { + fmt.Fprintf(os.Stderr, "error: %s requires a base URL\n", providerType) + os.Exit(1) + } } } + settingsJSON, err = providerSettingsJSON(settings) + if err != nil { + fmt.Fprintf(os.Stderr, "error: %v\n", err) + os.Exit(1) + } if !modelSet && model == "" { ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) selected, selectErr := promptProviderModelSelection(ctx, providerType, apiKey, baseURL, settings, scanner, os.Stdout, wfprovider.ListModelsWithSettings) @@ -893,6 +904,46 @@ func bedrockProviderSettings(accessKeyID, region string) (map[string]string, err return settings, nil } +func providerPromptsBaseURL(providerType string) bool { + switch providerType { + case "custom", "openai", "openai_compatible", "anthropic_compatible": + return true + default: + return false + } +} + +func providerRequiresBaseURL(providerType string) bool { + switch providerType { + case "custom", "openai_compatible", "anthropic_compatible": + return true + default: + return false + } +} + +func promptCustomProviderCompatibility(scanner *bufio.Scanner, out io.Writer) (map[string]string, error) { + fmt.Fprintln(out, "API compatibility:") + fmt.Fprintln(out, " 1. OpenAI-compatible /v1/chat/completions") + fmt.Fprintln(out, " 2. Anthropic-compatible /v1/messages") + fmt.Fprint(out, "Select [1]: ") + if !scanner.Scan() { + if scanErr := scanner.Err(); scanErr != nil { + return nil, scanErr + } + return map[string]string{"api_compat": "openai"}, nil + } + choice := strings.TrimSpace(scanner.Text()) + switch strings.ToLower(choice) { + case "", "1", "openai", "openai_compatible": + return map[string]string{"api_compat": "openai"}, nil + case "2", "anthropic", "anthropic_compatible": + return map[string]string{"api_compat": "anthropic"}, nil + default: + return nil, fmt.Errorf("invalid API compatibility selection %q", choice) + } +} + func providerSettingsJSON(settings map[string]string) (string, error) { if len(settings) == 0 { return "", nil diff --git a/cmd/ratchet/cmd_provider_test.go b/cmd/ratchet/cmd_provider_test.go index f951255..56e11e6 100644 --- a/cmd/ratchet/cmd_provider_test.go +++ b/cmd/ratchet/cmd_provider_test.go @@ -318,6 +318,47 @@ func TestPromptBedrockProviderCredentialsRequiresSecret(t *testing.T) { } } +func TestProviderBaseURLPromptPolicy(t *testing.T) { + for _, providerType := range []string{"custom", "openai", "openai_compatible", "anthropic_compatible"} { + if !providerPromptsBaseURL(providerType) { + t.Fatalf("%s should prompt for base URL", providerType) + } + } + for _, providerType := range []string{"custom", "openai_compatible", "anthropic_compatible"} { + if !providerRequiresBaseURL(providerType) { + t.Fatalf("%s should require base URL", providerType) + } + } + if providerRequiresBaseURL("openai") { + t.Fatal("openai should allow the default upstream URL") + } + if providerPromptsBaseURL("bedrock") { + t.Fatal("bedrock should use AWS region/settings rather than base URL by default") + } +} + +func TestPromptCustomProviderCompatibilityDefaultsOpenAI(t *testing.T) { + scanner := bufio.NewScanner(strings.NewReader("\n")) + settings, err := promptCustomProviderCompatibility(scanner, &strings.Builder{}) + if err != nil { + t.Fatalf("promptCustomProviderCompatibility: %v", err) + } + if settings["api_compat"] != "openai" { + t.Fatalf("settings = %#v", settings) + } +} + +func TestPromptCustomProviderCompatibilitySupportsAnthropic(t *testing.T) { + scanner := bufio.NewScanner(strings.NewReader("2\n")) + settings, err := promptCustomProviderCompatibility(scanner, &strings.Builder{}) + if err != nil { + t.Fatalf("promptCustomProviderCompatibility: %v", err) + } + if settings["api_compat"] != "anthropic" { + t.Fatalf("settings = %#v", settings) + } +} + func TestPromptProviderModelSelectionDefaultsToFirstEnumeratedModel(t *testing.T) { scanner := bufio.NewScanner(strings.NewReader("\n")) model, err := promptProviderModelSelection( diff --git a/go.mod b/go.mod index 318ad3f..c9e00fd 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.5 + github.com/GoCodeAlone/workflow-plugin-agent v0.12.6 github.com/charmbracelet/glamour v0.10.0 github.com/coder/acp-go-sdk v0.6.3 github.com/creack/pty v1.1.24 @@ -54,7 +54,7 @@ require ( 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.42.1 // indirect - github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.8 // indirect + github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.14 // 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 @@ -62,6 +62,7 @@ require ( 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/bedrockruntime v1.54.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 diff --git a/go.sum b/go.sum index 5fd1142..a6ed6a2 100644 --- a/go.sum +++ b/go.sum @@ -58,6 +58,8 @@ github.com/GoCodeAlone/workflow-plugin-acpx v0.1.0 h1:mwwaLAfsxEIYdIQ4wvV/PkT0Ol github.com/GoCodeAlone/workflow-plugin-acpx v0.1.0/go.mod h1:Bun12gpt+/WmqlAafEXpxSG1esbI3TAmwWFrgaR1xOY= 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-agent v0.12.6 h1:4to8WuvE6HGvTBO/samh675XZCm4QUUCzsiuH+3s1uA= +github.com/GoCodeAlone/workflow-plugin-agent v0.12.6/go.mod h1:1lLcDxpcL6duN/Ly8lxCmAdCeZAH12aSfSFowSPNwjg= 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= @@ -105,6 +107,8 @@ github.com/aws/aws-sdk-go-v2 v1.42.1 h1:9eOTgu1z/dVtYpNZ3/8/XbbaX0x/BqE3HUzAzs6K 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/aws/protocol/eventstream v1.7.14 h1:3IZY0XAJquT3aHzbkHfPzy4ACPcEjVG0x87KOwtpqGY= +github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.14/go.mod h1:zwM6veDkhGgQFqkBy+uT28AAYpLu+uFMlPl+rCg/73E= github.com/aws/aws-sdk-go-v2/config v1.32.16 h1:Q0iQ7quUgJP0F/SCRTieScnaMdXr9h/2+wze1u3cNeM= github.com/aws/aws-sdk-go-v2/config v1.32.16/go.mod h1:duCCnJEFqpt2RC6no1iK6q+8HpwOAkiUua0pY507dQc= github.com/aws/aws-sdk-go-v2/credentials v1.19.15 h1:fyvgWTszojq8hEnMi8PPBTvZdTtEVmAVyo+NFLHBhH4= @@ -119,6 +123,8 @@ github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.23 h1:FPXsW9+gMuIeKmz7j6ENWcWtBGT 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/bedrockruntime v1.54.2 h1:qmKlhMqcFouMkrntKDOZ93vDQLBAoVCVVNMo5sJmq8o= +github.com/aws/aws-sdk-go-v2/service/bedrockruntime v1.54.2/go.mod h1:RRUdkfdYMMT5wzMXS7pZ6JvsrW1e9XqJgKQq2ie3rIk= 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=