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
65 changes: 58 additions & 7 deletions cmd/ratchet/cmd_provider.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down
41 changes: 41 additions & 0 deletions cmd/ratchet/cmd_provider_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
5 changes: 3 additions & 2 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.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
Expand Down Expand Up @@ -54,14 +54,15 @@ 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
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/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
Expand Down
6 changes: 6 additions & 0 deletions go.sum
Original file line number Diff line number Diff line change
Expand Up @@ -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=
Expand Down Expand Up @@ -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=
Expand All @@ -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=
Expand Down
Loading