Skip to content
Open
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
16 changes: 8 additions & 8 deletions cmd/opencodereview/apply_provider_field_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@ func TestApplyProviderField(t *testing.T) {
{"extra_body", `{"k":1}`, func(e ProviderEntry) bool { return e.ExtraBody["k"] != nil }},
}
for _, c := range cases {
if err := applyProviderField(&e, c.field, "providers.p."+c.field, c.value); err != nil {
if err := applyProviderField("p", &e, c.field, "providers.p."+c.field, c.value); err != nil {
t.Fatalf("field %q: %v", c.field, err)
}
if !c.check(e) {
Expand All @@ -34,20 +34,20 @@ func TestApplyProviderField(t *testing.T) {

t.Run("protocol validated and normalized", func(t *testing.T) {
var e ProviderEntry
if err := applyProviderField(&e, "protocol", "providers.p.protocol", "openai"); err != nil {
if err := applyProviderField("p", &e, "protocol", "providers.p.protocol", "openai"); err != nil {
t.Fatalf("valid protocol: %v", err)
}
if e.Protocol == "" {
t.Error("protocol not set")
}
if err := applyProviderField(&e, "protocol", "providers.p.protocol", "not-a-protocol"); err == nil {
if err := applyProviderField("p", &e, "protocol", "providers.p.protocol", "not-a-protocol"); err == nil {
t.Error("expected error for invalid protocol")
}
})

t.Run("auth_header normalized", func(t *testing.T) {
var e ProviderEntry
if err := applyProviderField(&e, "auth_header", "providers.p.auth_header", "x-api-key"); err != nil {
if err := applyProviderField("p", &e, "auth_header", "providers.p.auth_header", "x-api-key"); err != nil {
t.Fatalf("valid auth header: %v", err)
}
if e.AuthHeader == "" {
Expand All @@ -57,21 +57,21 @@ func TestApplyProviderField(t *testing.T) {

t.Run("auth_header rejects unsupported value", func(t *testing.T) {
var e ProviderEntry
if err := applyProviderField(&e, "auth_header", "providers.p.auth_header", "cookie"); err == nil {
if err := applyProviderField("p", &e, "auth_header", "providers.p.auth_header", "cookie"); err == nil {
t.Error("expected error for unsupported auth header")
}
})

t.Run("extra_body rejects invalid JSON", func(t *testing.T) {
var e ProviderEntry
if err := applyProviderField(&e, "extra_body", "providers.p.extra_body", "{bad"); err == nil {
if err := applyProviderField("p", &e, "extra_body", "providers.p.extra_body", "{bad"); err == nil {
t.Error("expected JSON error")
}
})

t.Run("extra_headers parsed", func(t *testing.T) {
var e ProviderEntry
if err := applyProviderField(&e, "extra_headers", "providers.p.extra_headers", "X-A=1"); err != nil {
if err := applyProviderField("p", &e, "extra_headers", "providers.p.extra_headers", "X-A=1"); err != nil {
t.Fatalf("valid extra headers: %v", err)
}
if len(e.ExtraHeaders) == 0 {
Expand All @@ -81,7 +81,7 @@ func TestApplyProviderField(t *testing.T) {

t.Run("unknown field returns error", func(t *testing.T) {
var e ProviderEntry
if err := applyProviderField(&e, "bogus", "providers.p.bogus", "x"); err == nil {
if err := applyProviderField("p", &e, "bogus", "providers.p.bogus", "x"); err == nil {
t.Error("expected error for unknown field")
}
})
Expand Down
256 changes: 256 additions & 0 deletions cmd/opencodereview/bedrock_config_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,256 @@
// SPDX-License-Identifier: Apache-2.0
// Copyright 2026 alibaba/open-code-review Contributors

package main

import (
"encoding/json"
"os"
"path/filepath"
"strings"
"testing"

"github.com/alibaba/open-code-review/internal/llm"
)

// TestConfigRoundTripKeepsAWSSettings is the regression test for a silent loss:
// config is unmarshalled into Config and marshalled back on every write, so
// before aws_profile / aws_region existed on ProviderEntry, the first run of any
// config command deleted them from a hand-written file — with no error, and no
// way for the user to tell why Bedrock suddenly used the wrong region.
func TestConfigRoundTripKeepsAWSSettings(t *testing.T) {
path := filepath.Join(t.TempDir(), "config.json")
original := `{
"provider": "bedrock",
"model": "us.anthropic.claude-sonnet-4-6",
"providers": {
"bedrock": { "aws_region": "us-west-2", "aws_profile": "example-profile" }
}
}`
if err := os.WriteFile(path, []byte(original), 0o600); err != nil {
t.Fatalf("write config: %v", err)
}

cfg, err := loadOrCreateConfig(path)
if err != nil {
t.Fatalf("loadOrCreateConfig: %v", err)
}
if err := saveConfig(path, cfg); err != nil {
t.Fatalf("saveConfig: %v", err)
}

reloaded, err := loadOrCreateConfig(path)
if err != nil {
t.Fatalf("reload: %v", err)
}
entry := reloaded.Providers["bedrock"]
if entry.AWSRegion != "us-west-2" {
t.Errorf("AWSRegion = %q after round trip, want us-west-2", entry.AWSRegion)
}
if entry.AWSProfile != "example-profile" {
t.Errorf("AWSProfile = %q after round trip, want example-profile", entry.AWSProfile)
}

// The resolver reads the same file independently; assert the written JSON
// still carries the keys it looks for, not just that our struct held them.
data, err := os.ReadFile(path)
if err != nil {
t.Fatalf("read back: %v", err)
}
var raw map[string]any
if err := json.Unmarshal(data, &raw); err != nil {
t.Fatalf("unmarshal written config: %v", err)
}
providers, _ := raw["providers"].(map[string]any)
bedrockEntry, _ := providers["bedrock"].(map[string]any)
if bedrockEntry["aws_region"] != "us-west-2" || bedrockEntry["aws_profile"] != "example-profile" {
t.Errorf("written JSON = %v, want aws_region and aws_profile preserved", bedrockEntry)
}
}

func TestSetProviderValueAWSSettings(t *testing.T) {
tests := []struct {
name string
key string
value string
wantErr string
check func(*testing.T, *Config)
}{
{
name: "region on an ambient provider",
key: "providers.bedrock.aws_region",
value: "us-west-2",
check: func(t *testing.T, cfg *Config) {
if got := cfg.Providers["bedrock"].AWSRegion; got != "us-west-2" {
t.Errorf("AWSRegion = %q, want us-west-2", got)
}
},
},
{
name: "profile is trimmed",
key: "providers.bedrock.aws_profile",
value: " example-profile ",
check: func(t *testing.T, cfg *Config) {
if got := cfg.Providers["bedrock"].AWSProfile; got != "example-profile" {
t.Errorf("AWSProfile = %q, want example-profile", got)
}
},
},
{
name: "empty value hands the decision back to the AWS chain",
key: "providers.bedrock.aws_profile",
value: "",
check: func(t *testing.T, cfg *Config) {
if got := cfg.Providers["bedrock"].AWSProfile; got != "" {
t.Errorf("AWSProfile = %q, want empty", got)
}
},
},
{
// Storing it would be dead config that reads as applied.
name: "rejected on a key-based provider",
key: "providers.anthropic.aws_region",
value: "us-west-2",
wantErr: "does not apply to provider",
},
{
name: "whitespace inside the value is rejected",
key: "providers.bedrock.aws_region",
value: "us west 2",
wantErr: "contains whitespace",
},
}

for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
cfg := &Config{}
err := setProviderValue(cfg, tc.key, tc.value)
if tc.wantErr != "" {
if err == nil {
t.Fatalf("setProviderValue(%q, %q) = nil, want error containing %q", tc.key, tc.value, tc.wantErr)
}
if !strings.Contains(err.Error(), tc.wantErr) {
t.Fatalf("error = %q, want it to contain %q", err, tc.wantErr)
}
return
}
if err != nil {
t.Fatalf("setProviderValue(%q, %q): %v", tc.key, tc.value, err)
}
tc.check(t, cfg)
})
}
}

// TestSetCustomProviderAWSSettingsFollowProtocol covers the custom-provider
// path: aws_* is meaningful there only once the entry speaks the Bedrock
// protocol, so the order of the two set commands matters and the error has to
// say why.
func TestSetCustomProviderAWSSettingsFollowProtocol(t *testing.T) {
cfg := &Config{}
if err := setCustomProviderValue(cfg, "custom_providers.mine.aws_region", "us-west-2"); err == nil {
t.Fatal("aws_region accepted before a protocol was set; want an error")
}

if err := setCustomProviderValue(cfg, "custom_providers.mine.protocol", llm.ProtocolAnthropicBedrock); err != nil {
t.Fatalf("set protocol: %v", err)
}
if err := setCustomProviderValue(cfg, "custom_providers.mine.aws_region", "us-west-2"); err != nil {
t.Fatalf("set aws_region after protocol: %v", err)
}
if got := cfg.CustomProviders["mine"].AWSRegion; got != "us-west-2" {
t.Errorf("AWSRegion = %q, want us-west-2", got)
}
}

// TestAWSSettingsRejectedWhenEntryOverridesProtocol covers the same
// entry-level protocol override the resolver honours: a preset's protocol can be
// overridden per entry, so `protocol: openai` on the bedrock preset must stop
// accepting AWS settings that nothing would read.
func TestAWSSettingsRejectedWhenEntryOverridesProtocol(t *testing.T) {
cfg := &Config{}
if err := setProviderValue(cfg, "providers.bedrock.protocol", "openai"); err != nil {
t.Fatalf("set protocol: %v", err)
}
err := setProviderValue(cfg, "providers.bedrock.aws_region", "us-west-2")
if err == nil {
t.Fatal("aws_region accepted on a bedrock entry overridden to protocol openai; want an error")
}
if !strings.Contains(err.Error(), "does not apply to provider") {
t.Errorf("error = %q, want it to explain the field does not apply", err)
}

// Overriding back to the bedrock protocol makes them meaningful again.
if err := setProviderValue(cfg, "providers.bedrock.protocol", llm.ProtocolAnthropicBedrock); err != nil {
t.Fatalf("set protocol back: %v", err)
}
if err := setProviderValue(cfg, "providers.bedrock.aws_region", "us-west-2"); err != nil {
t.Errorf("aws_region rejected for an explicit bedrock protocol: %v", err)
}
}

func TestCheckAPIKeyRequirement(t *testing.T) {
bedrock, ok := llm.LookupProvider("bedrock")
if !ok {
t.Fatal("bedrock preset not registered")
}
anthropic, ok := llm.LookupProvider("anthropic")
if !ok {
t.Fatal("anthropic preset not registered")
}

if err := checkAPIKeyRequirement("bedrock", "", bedrock, true); err != nil {
t.Errorf("ambient provider with no api_key = %v, want nil", err)
}

t.Setenv(anthropic.EnvVar, "")
if err := checkAPIKeyRequirement("anthropic", "", anthropic, true); err == nil {
t.Error("key-based provider with no api_key and no env var = nil, want an error")
}
}

// TestProviderTUIAmbientProviderSkipsAPIKeyStep pins the wizard flow: the model
// step is the last one for a provider with no key to collect. An API-key prompt
// that must be left blank reads as a step the user failed to complete.
func TestProviderTUIAmbientProviderSkipsAPIKeyStep(t *testing.T) {
m := newProviderTUI(&Config{}, "")
idx := -1
for i, p := range m.providers {
if p.Name == "bedrock" {
idx = i
break
}
}
if idx < 0 {
t.Fatal("bedrock not offered in the official provider list")
}
m.officialIdx = idx

result, _ := m.Update(enterKey())
atModel := result.(providerTUIModel)
if atModel.step != stepModel {
t.Fatalf("after Enter on provider, step = %d, want %d (stepModel)", atModel.step, stepModel)
}

result, cmd := atModel.Update(enterKey())
done := result.(providerTUIModel)
if done.step == stepAPIKey {
t.Error("ambient provider advanced to stepAPIKey; want the model step to be final")
}
if !done.confirmed {
t.Error("confirmed = false; want the selection confirmed from the model step")
}
if cmd == nil {
t.Error("no command returned; want tea.Quit")
}
res := done.result()
if res.provider != "bedrock" {
t.Errorf("result provider = %q, want bedrock", res.provider)
}
if res.apiKey != "" {
t.Errorf("result apiKey = %q, want empty for an ambient provider", res.apiKey)
}
if got := res.resolvedModel(); got == "" {
t.Error("resolvedModel is empty; want the model selected on the model step")
}
}
Loading