diff --git a/aibridge/aibridgetest/aibridgetest.go b/aibridge/aibridgetest/aibridgetest.go new file mode 100644 index 0000000000..c86bee683a --- /dev/null +++ b/aibridge/aibridgetest/aibridgetest.go @@ -0,0 +1,19 @@ +package aibridgetest + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/aibridge" +) + +// NewAnthropicProvider builds an Anthropic provider for tests, failing the test +// if credential resolution fails. +func NewAnthropicProvider(t testing.TB, cfg aibridge.AnthropicConfig, bedrockCfg *aibridge.AWSBedrockConfig) aibridge.Provider { + t.Helper() + p, err := aibridge.NewAnthropicProvider(context.Background(), cfg, bedrockCfg) + require.NoError(t, err) + return p +} diff --git a/aibridge/api.go b/aibridge/api.go index 34dce84ef8..b74912d565 100644 --- a/aibridge/api.go +++ b/aibridge/api.go @@ -45,8 +45,8 @@ func AsActor(ctx context.Context, actorID string, metadata recorder.Metadata) co return aibcontext.AsActor(ctx, actorID, metadata) } -func NewAnthropicProvider(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) provider.Provider { - return provider.NewAnthropic(cfg, bedrockCfg) +func NewAnthropicProvider(ctx context.Context, cfg config.Anthropic, bedrockCfg *config.AWSBedrock) (provider.Provider, error) { + return provider.NewAnthropic(ctx, cfg, bedrockCfg) } func NewOpenAIProvider(cfg config.OpenAI) provider.Provider { diff --git a/aibridge/bridge_test.go b/aibridge/bridge_test.go index 9ac7ea9ec3..76242f20be 100644 --- a/aibridge/bridge_test.go +++ b/aibridge/bridge_test.go @@ -15,6 +15,7 @@ import ( "cdr.dev/slog/v3/sloggers/slogtest" "github.com/coder/coder/v2/aibridge" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/config" "github.com/coder/coder/v2/aibridge/internal/testutil" "github.com/coder/coder/v2/aibridge/provider" @@ -36,7 +37,7 @@ func TestValidateProviders(t *testing.T) { name: "all_supported_providers", providers: []provider.Provider{ aibridge.NewOpenAIProvider(config.OpenAI{Name: "openai", BaseURL: "https://api.openai.com/v1/"}), - aibridge.NewAnthropicProvider(config.Anthropic{Name: "anthropic", BaseURL: "https://api.anthropic.com/"}, nil), + aibridgetest.NewAnthropicProvider(t, config.Anthropic{Name: "anthropic", BaseURL: "https://api.anthropic.com/"}, nil), aibridge.NewCopilotProvider(config.Copilot{Name: "copilot", BaseURL: "https://api.individual.githubcopilot.com"}), aibridge.NewCopilotProvider(config.Copilot{Name: "copilot-business", BaseURL: "https://api.business.githubcopilot.com"}), aibridge.NewCopilotProvider(config.Copilot{Name: "copilot-enterprise", BaseURL: "https://api.enterprise.githubcopilot.com"}), @@ -46,7 +47,7 @@ func TestValidateProviders(t *testing.T) { name: "default_names_and_base_urls", providers: []provider.Provider{ aibridge.NewOpenAIProvider(config.OpenAI{}), - aibridge.NewAnthropicProvider(config.Anthropic{}, nil), + aibridgetest.NewAnthropicProvider(t, config.Anthropic{}, nil), aibridge.NewCopilotProvider(config.Copilot{}), }, }, @@ -126,13 +127,13 @@ func TestPassthroughRoutesForProviders(t *testing.T) { name string baseURLPath string requestPath string - provider func(string) provider.Provider + provider func(*testing.T, string) provider.Provider expectPath string }{ { name: "openAI_no_base_path", requestPath: "/openai/v1/conversations", - provider: func(baseURL string) provider.Provider { + provider: func(_ *testing.T, baseURL string) provider.Provider { return aibridge.NewOpenAIProvider(config.OpenAI{BaseURL: baseURL}) }, expectPath: "/conversations", @@ -141,7 +142,7 @@ func TestPassthroughRoutesForProviders(t *testing.T) { name: "openAI_with_base_path", baseURLPath: "/v1", requestPath: "/openai/v1/conversations", - provider: func(baseURL string) provider.Provider { + provider: func(_ *testing.T, baseURL string) provider.Provider { return aibridge.NewOpenAIProvider(config.OpenAI{BaseURL: baseURL}) }, expectPath: "/v1/conversations", @@ -149,8 +150,8 @@ func TestPassthroughRoutesForProviders(t *testing.T) { { name: "anthropic_no_base_path", requestPath: "/anthropic/v1/models", - provider: func(baseURL string) provider.Provider { - return aibridge.NewAnthropicProvider(config.Anthropic{BaseURL: baseURL}, nil) + provider: func(t *testing.T, baseURL string) provider.Provider { + return aibridgetest.NewAnthropicProvider(t, config.Anthropic{BaseURL: baseURL}, nil) }, expectPath: "/v1/models", }, @@ -158,15 +159,15 @@ func TestPassthroughRoutesForProviders(t *testing.T) { name: "anthropic_with_base_path", baseURLPath: "/v1", requestPath: "/anthropic/v1/models", - provider: func(baseURL string) provider.Provider { - return aibridge.NewAnthropicProvider(config.Anthropic{BaseURL: baseURL}, nil) + provider: func(t *testing.T, baseURL string) provider.Provider { + return aibridgetest.NewAnthropicProvider(t, config.Anthropic{BaseURL: baseURL}, nil) }, expectPath: "/v1/v1/models", }, { name: "copilot_no_base_path", requestPath: "/copilot/models", - provider: func(baseURL string) provider.Provider { + provider: func(_ *testing.T, baseURL string) provider.Provider { return aibridge.NewCopilotProvider(config.Copilot{BaseURL: baseURL}) }, expectPath: "/models", @@ -175,7 +176,7 @@ func TestPassthroughRoutesForProviders(t *testing.T) { name: "copilot_with_base_path", baseURLPath: "/v1", requestPath: "/copilot/models", - provider: func(baseURL string) provider.Provider { + provider: func(_ *testing.T, baseURL string) provider.Provider { return aibridge.NewCopilotProvider(config.Copilot{BaseURL: baseURL}) }, expectPath: "/v1/models", @@ -196,7 +197,7 @@ func TestPassthroughRoutesForProviders(t *testing.T) { t.Cleanup(upstream.Close) rec := testutil.MockRecorder{} - prov := tc.provider(upstream.URL + tc.baseURLPath) + prov := tc.provider(t, upstream.URL+tc.baseURLPath) bridge, err := aibridge.NewRequestBridge(t.Context(), []provider.Provider{prov}, &rec, nil, logger, nil, bridgeTestTracer) require.NoError(t, err) @@ -213,13 +214,13 @@ func TestPassthroughRoutesForProviders(t *testing.T) { func TestRequestBodySizeLimit(t *testing.T) { t.Parallel() - newOpenAI := func(baseURL string) provider.Provider { + newOpenAI := func(_ *testing.T, baseURL string) provider.Provider { return aibridge.NewOpenAIProvider(config.OpenAI{Name: "openai", BaseURL: baseURL}) } - newAnthropic := func(baseURL string) provider.Provider { - return aibridge.NewAnthropicProvider(config.Anthropic{Name: "anthropic", BaseURL: baseURL}, nil) + newAnthropic := func(t *testing.T, baseURL string) provider.Provider { + return aibridgetest.NewAnthropicProvider(t, config.Anthropic{Name: "anthropic", BaseURL: baseURL}, nil) } - newCopilot := func(baseURL string) provider.Provider { + newCopilot := func(_ *testing.T, baseURL string) provider.Provider { return aibridge.NewCopilotProvider(config.Copilot{Name: "copilot", BaseURL: baseURL}) } @@ -232,7 +233,7 @@ func TestRequestBodySizeLimit(t *testing.T) { tests := []struct { name string - provider func(baseURL string) provider.Provider + provider func(*testing.T, string) provider.Provider path string body []byte }{ @@ -258,7 +259,7 @@ func TestRequestBodySizeLimit(t *testing.T) { })) t.Cleanup(upstream.Close) - prov := tc.provider(upstream.URL) + prov := tc.provider(t, upstream.URL) bridge, err := aibridge.NewRequestBridge( t.Context(), []provider.Provider{prov}, diff --git a/aibridge/config/config.go b/aibridge/config/config.go index 9fdeb32d57..3dc76841aa 100644 --- a/aibridge/config/config.go +++ b/aibridge/config/config.go @@ -33,6 +33,11 @@ type AWSBedrock struct { // (https://bedrock-runtime.{region}.amazonaws.com). // This is useful for routing requests through a proxy or for testing. BaseURL string + // RoleARN, when set, is assumed via STS before calling Bedrock. The base + // identity (static keys or the AWS SDK default credential chain, e.g. + // IRSA / EKS Pod Identity / EC2 Instance Profile) signs the AssumeRole + // call, and the resulting temporary credentials sign Bedrock requests. + RoleARN string } // OpenAI carries configuration for an OpenAI provider. diff --git a/aibridge/intercept/credential.go b/aibridge/intercept/credential.go index 008c43463d..24016b6950 100644 --- a/aibridge/intercept/credential.go +++ b/aibridge/intercept/credential.go @@ -28,7 +28,7 @@ const ( // before failover selects a key, and a key resolved dynamically at request time. const ( hintFailoverKey = "" - hintBedrockChainKey = "" + hintBedrockChainKey = "" ) // Credential is the per-request upstream authentication for an interception: diff --git a/aibridge/intercept/credential_test.go b/aibridge/intercept/credential_test.go index 0f78d8dee0..84e27e083a 100644 --- a/aibridge/intercept/credential_test.go +++ b/aibridge/intercept/credential_test.go @@ -19,6 +19,9 @@ import ( func TestCredential(t *testing.T) { t.Parallel() + // Matches the VARCHAR(15) DB constraint. + const maxCredentialHintLength = 15 + tests := []struct { name string newCred func(t *testing.T) intercept.Credential @@ -72,7 +75,7 @@ func TestCredential(t *testing.T) { }, expectKind: intercept.CredentialKindCentralized, expectAuthHeader: "", - expectHint: "", + expectHint: "", expectLength: 0, }, { @@ -118,6 +121,8 @@ func TestCredential(t *testing.T) { assert.Equal(t, tc.expectKind, cred.Kind(), "Kind") assert.Equal(t, tc.expectAuthHeader, cred.AuthHeader(), "AuthHeader") assert.Equal(t, tc.expectHint, cred.Hint(), "Hint") + assert.LessOrEqual(t, len(cred.Hint()), maxCredentialHintLength, + "Hint must fit the credential_hint column") assert.Equal(t, tc.expectLength, cred.Length(), "Length") credBYOK, credBYOKOK := intercept.AsBYOK(cred) diff --git a/aibridge/intercept/messages/base.go b/aibridge/intercept/messages/base.go index dff58c1cad..7c895595e9 100644 --- a/aibridge/intercept/messages/base.go +++ b/aibridge/intercept/messages/base.go @@ -16,8 +16,7 @@ import ( "github.com/anthropics/anthropic-sdk-go/option" "github.com/anthropics/anthropic-sdk-go/shared" "github.com/anthropics/anthropic-sdk-go/shared/constant" - "github.com/aws/aws-sdk-go-v2/config" - "github.com/aws/aws-sdk-go-v2/credentials" + "github.com/aws/aws-sdk-go-v2/aws" "github.com/google/uuid" "go.opentelemetry.io/otel/attribute" "go.opentelemetry.io/otel/trace" @@ -65,13 +64,26 @@ var bedrockSupportedBetaFlags = map[string]bool{ "tool-examples-2025-10-29": true, } +// BedrockRuntime carries everything a Bedrock-backed interception needs: the +// static Bedrock config plus the AWS credentials provider. +type BedrockRuntime struct { + Cfg aibconfig.AWSBedrock + Creds aws.CredentialsProvider + // ResolvedRegion is the region the AWS SDK resolved at construction (from + // the environment, shared config, or IMDS). It is used for request signing + // when Cfg.Region is empty, e.g. a custom base URL with the region supplied + // via AWS_REGION. + ResolvedRegion string +} + type interceptionBase struct { id uuid.UUID reqPayload RequestPayload - cfg intercept.Config - cred intercept.Credential - bedrockCfg *aibconfig.AWSBedrock + cfg intercept.Config + cred intercept.Credential + // bedrock is nil for non-Bedrock providers. + bedrock *BedrockRuntime // clientHeaders are the original HTTP headers from the client request. clientHeaders http.Header @@ -107,10 +119,10 @@ func (i *interceptionBase) Model() string { return "coder-aibridge-unknown" } - if i.bedrockCfg != nil { - model := i.bedrockCfg.Model + if i.bedrock != nil { + model := i.bedrock.Cfg.Model if i.isSmallFastModel() { - model = i.bedrockCfg.SmallFastModel + model = i.bedrock.Cfg.SmallFastModel } return model } @@ -126,7 +138,7 @@ func (i *interceptionBase) baseTraceAttributes(r *http.Request, streaming bool) attribute.String(tracing.Provider, i.cfg.ProviderName), attribute.String(tracing.Model, i.Model()), attribute.Bool(tracing.Streaming, streaming), - attribute.Bool(tracing.IsBedrock, i.bedrockCfg != nil), + attribute.Bool(tracing.IsBedrock, i.bedrock != nil), } } @@ -238,10 +250,10 @@ func (i *interceptionBase) newMessagesService(ctx context.Context, opts ...optio opts = append(opts, option.WithMiddleware(mw)) } - if i.bedrockCfg != nil { + if i.bedrock != nil { ctx, cancel := context.WithTimeout(ctx, time.Second*30) defer cancel() - bedrockOpts, err := i.withAWSBedrockOptions(ctx, i.bedrockCfg) + bedrockOpts, err := i.withAWSBedrockOptions(ctx) if err != nil { return anthropic.MessageService{}, err } @@ -262,14 +274,13 @@ func (i *interceptionBase) withBody() option.RequestOption { // withAWSBedrockOptions returns request options for authenticating with AWS Bedrock. // -// When both AccessKey and AccessKeySecret are set in the aibridge config, they are -// used directly as static credentials. Otherwise, the AWS SDK default credential chain -// resolves credentials (environment variables, shared config/credentials files, IAM -// roles, IRSA, SSO, IMDS, etc.). -func (*interceptionBase) withAWSBedrockOptions(ctx context.Context, cfg *aibconfig.AWSBedrock) ([]option.RequestOption, error) { - if cfg == nil { - return nil, xerrors.New("nil config given") +// Credentials come from i.bedrock.Creds. It is a shared credentials cache, so the per-request Retrieve() +// below is served from that cache and does not re-resolve or re-assume on every request. +func (i *interceptionBase) withAWSBedrockOptions(ctx context.Context) ([]option.RequestOption, error) { + if i.bedrock == nil { + return nil, xerrors.New("nil bedrock runtime") } + cfg := i.bedrock.Cfg if cfg.Region == "" && cfg.BaseURL == "" { return nil, xerrors.New("region or base url required") } @@ -280,38 +291,22 @@ func (*interceptionBase) withAWSBedrockOptions(ctx context.Context, cfg *aibconf return nil, xerrors.New("small fast model required") } - loadOpts := []func(*config.LoadOptions) error{ - config.WithRegion(cfg.Region), + // Fail fast: ensure credentials can be resolved before signing. Served from + // the shared cache on most requests (no network); on the cold or refresh + // path this performs the actual STS/IMDS call. + if _, err := i.bedrock.Creds.Retrieve(ctx); err != nil { + return nil, xerrors.Errorf("resolve AWS credentials: %w", err) } - // Use static credentials when explicitly provided, otherwise fall back to the SDK default credential chain. - switch { - // Both set: use static credentials directly. - case cfg.AccessKey != "" && cfg.AccessKeySecret != "": - loadOpts = append(loadOpts, config.WithCredentialsProvider( - credentials.NewStaticCredentialsProvider( - cfg.AccessKey, - cfg.AccessKeySecret, - "", - ), - )) - // Only one set: misconfiguration. - case cfg.AccessKey != "" || cfg.AccessKeySecret != "": - return nil, xerrors.New("both access key and access key secret must be provided together") - // Neither set: SDK default credential chain resolves credentials. - default: + // Fall back to the SDK-resolved region (e.g. from AWS_REGION) when no + // explicit region is configured. + region := cfg.Region + if region == "" { + region = i.bedrock.ResolvedRegion } - - awsCfg, err := config.LoadDefaultConfig(ctx, loadOpts...) - if err != nil { - return nil, xerrors.Errorf("failed to load AWS Bedrock config: %w", err) - } - - // Fail fast: ensure credentials can be resolved before making any requests. - // awsCfg already carries the credentials provider, and the Bedrock middleware - // will call Retrieve on it when signing each request. - if _, err := awsCfg.Credentials.Retrieve(ctx); err != nil { - return nil, xerrors.Errorf("no AWS credentials found: %w", err) + awsCfg := aws.Config{ + Region: region, + Credentials: i.bedrock.Creds, } var out []option.RequestOption @@ -336,7 +331,7 @@ func (*interceptionBase) withAWSBedrockOptions(ctx context.Context, cfg *aibconf // don't support adaptive thinking natively, or enabled thinking to adaptive for models that only support // adaptive (Opus 4.7+). func (i *interceptionBase) augmentRequestForBedrock() { - if i.bedrockCfg == nil { + if i.bedrock == nil { return } diff --git a/aibridge/intercept/messages/base_internal_test.go b/aibridge/intercept/messages/base_internal_test.go index 5a4c5e6f30..0b2ee4d77b 100644 --- a/aibridge/intercept/messages/base_internal_test.go +++ b/aibridge/intercept/messages/base_internal_test.go @@ -9,6 +9,7 @@ import ( "github.com/anthropics/anthropic-sdk-go" "github.com/anthropics/anthropic-sdk-go/shared/constant" + "github.com/aws/aws-sdk-go-v2/credentials" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/tidwall/gjson" @@ -87,14 +88,14 @@ func TestAWSBedrockValidation(t *testing.T) { tests := []struct { name string - cfg *config.AWSBedrock + cfg config.AWSBedrock expectError bool errorMsg string }{ // Valid cases: static credentials. { name: "static credentials with region", - cfg: &config.AWSBedrock{ + cfg: config.AWSBedrock{ Region: "us-east-1", AccessKey: "test-key", AccessKeySecret: "test-secret", @@ -104,7 +105,7 @@ func TestAWSBedrockValidation(t *testing.T) { }, { name: "static credentials with base url", - cfg: &config.AWSBedrock{ + cfg: config.AWSBedrock{ BaseURL: "http://bedrock.internal", AccessKey: "test-key", AccessKeySecret: "test-secret", @@ -119,7 +120,7 @@ func TestAWSBedrockValidation(t *testing.T) { // // See TestAWSBedrockIntegration which validates this. name: "static credentials with base url & region", - cfg: &config.AWSBedrock{ + cfg: config.AWSBedrock{ Region: "us-east-1", AccessKey: "test-key", AccessKeySecret: "test-secret", @@ -130,7 +131,7 @@ func TestAWSBedrockValidation(t *testing.T) { // Invalid cases. { name: "missing region & base url", - cfg: &config.AWSBedrock{ + cfg: config.AWSBedrock{ Region: "", AccessKey: "test-key", AccessKeySecret: "test-secret", @@ -140,32 +141,9 @@ func TestAWSBedrockValidation(t *testing.T) { expectError: true, errorMsg: "region or base url required", }, - { - name: "missing access key", - cfg: &config.AWSBedrock{ - Region: "us-east-1", - AccessKeySecret: "test-secret", - Model: "test-model", - SmallFastModel: "test-small-model", - }, - expectError: true, - errorMsg: "both access key and access key secret must be provided together", - }, - { - name: "missing access key secret", - cfg: &config.AWSBedrock{ - Region: "us-east-1", - AccessKey: "test-key", - AccessKeySecret: "", - Model: "test-model", - SmallFastModel: "test-small-model", - }, - expectError: true, - errorMsg: "both access key and access key secret must be provided together", - }, { name: "missing model", - cfg: &config.AWSBedrock{ + cfg: config.AWSBedrock{ Region: "us-east-1", AccessKey: "test-key", AccessKeySecret: "test-secret", @@ -177,7 +155,7 @@ func TestAWSBedrockValidation(t *testing.T) { }, { name: "missing small fast model", - cfg: &config.AWSBedrock{ + cfg: config.AWSBedrock{ Region: "us-east-1", AccessKey: "test-key", AccessKeySecret: "test-secret", @@ -189,24 +167,23 @@ func TestAWSBedrockValidation(t *testing.T) { }, { name: "all fields empty", - cfg: &config.AWSBedrock{}, + cfg: config.AWSBedrock{}, expectError: true, errorMsg: "region or base url required", }, - { - name: "nil config", - cfg: nil, - expectError: true, - errorMsg: "nil config given", - }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { t.Parallel() - base := &interceptionBase{} - opts, err := base.withAWSBedrockOptions(context.Background(), tt.cfg) + base := &interceptionBase{ + bedrock: &BedrockRuntime{ + Cfg: tt.cfg, + Creds: credentials.NewStaticCredentialsProvider("test-key", "test-secret", ""), + }, + } + opts, err := base.withAWSBedrockOptions(context.Background()) if tt.expectError { require.Error(t, err) @@ -219,87 +196,16 @@ func TestAWSBedrockValidation(t *testing.T) { } } -// TestAWSBedrockCredentialChain tests credential resolution via the AWS SDK default credential chain. -// NOTE: Cannot use t.Parallel() here because subtests use t.Setenv which requires sequential execution. -func TestAWSBedrockCredentialChain(t *testing.T) { - tests := []struct { - name string - cfg *config.AWSBedrock - envVars map[string]string - expectError bool - errorMsg string - }{ - { - name: "temporary credentials via env", - cfg: &config.AWSBedrock{ - Region: "us-east-1", - Model: "test-model", - SmallFastModel: "test-small-model", - }, - envVars: map[string]string{ - "AWS_ACCESS_KEY_ID": "test-key", - "AWS_SECRET_ACCESS_KEY": "test-secret", - }, - }, - { - name: "temporary credentials with session token via env", - cfg: &config.AWSBedrock{ - Region: "us-east-1", - Model: "test-model", - SmallFastModel: "test-small-model", - }, - envVars: map[string]string{ - "AWS_ACCESS_KEY_ID": "test-key", - "AWS_SECRET_ACCESS_KEY": "test-secret", - "AWS_SESSION_TOKEN": "test-session-token", - }, - }, - { - // When static credentials are not provided and no environment credentials are set, - // the SDK default credential chain fails to resolve credentials. - name: "error when no credential source is configured", - cfg: &config.AWSBedrock{ - Region: "us-east-1", - Model: "test-model", - SmallFastModel: "test-small-model", - }, - envVars: map[string]string{ - "AWS_ACCESS_KEY_ID": "", - "AWS_SECRET_ACCESS_KEY": "", - "AWS_SESSION_TOKEN": "", - "AWS_PROFILE": "", - "AWS_SHARED_CREDENTIALS_FILE": "/dev/null", - "AWS_CONFIG_FILE": "/dev/null", - "AWS_WEB_IDENTITY_TOKEN_FILE": "", - "AWS_ROLE_ARN": "", - "AWS_ROLE_SESSION_NAME": "", - "AWS_CONTAINER_CREDENTIALS_RELATIVE_URI": "", - "AWS_CONTAINER_CREDENTIALS_FULL_URI": "", - "AWS_CONTAINER_AUTHORIZATION_TOKEN": "", - "AWS_EC2_METADATA_DISABLED": "true", - }, - expectError: true, - errorMsg: "no AWS credentials found", - }, - } +// TestAWSBedrockOptionsRequireRuntime verifies that option assembly fails when +// the Bedrock runtime was not set. This should never happen in practice, since +// withAWSBedrockOptions is only called when i.bedrock != nil. +func TestAWSBedrockOptionsRequireRuntime(t *testing.T) { + t.Parallel() - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - for key, val := range tt.envVars { - t.Setenv(key, val) - } - base := &interceptionBase{} - opts, err := base.withAWSBedrockOptions(context.Background(), tt.cfg) - - if tt.expectError { - require.Error(t, err) - require.Contains(t, err.Error(), tt.errorMsg) - } else { - require.NotEmpty(t, opts) - require.NoError(t, err) - } - }) - } + base := &interceptionBase{} + _, err := base.withAWSBedrockOptions(context.Background()) + require.Error(t, err) + require.Contains(t, err.Error(), "nil bedrock runtime") } func TestAccumulateUsage(t *testing.T) { @@ -878,9 +784,11 @@ func TestAugmentRequestForBedrock_AdaptiveThinking(t *testing.T) { i := &interceptionBase{ reqPayload: mustMessagesPayload(t, tc.requestBody), - bedrockCfg: &config.AWSBedrock{ - Model: tc.bedrockModel, - SmallFastModel: "anthropic.claude-haiku-3-5", + bedrock: &BedrockRuntime{ + Cfg: config.AWSBedrock{ + Model: tc.bedrockModel, + SmallFastModel: "anthropic.claude-haiku-3-5", + }, }, clientHeaders: clientHeaders, logger: slog.Make(), diff --git a/aibridge/intercept/messages/blocking.go b/aibridge/intercept/messages/blocking.go index e6b5a8a596..8b3d9cc613 100644 --- a/aibridge/intercept/messages/blocking.go +++ b/aibridge/intercept/messages/blocking.go @@ -17,7 +17,6 @@ import ( "golang.org/x/xerrors" "cdr.dev/slog/v3" - aibconfig "github.com/coder/coder/v2/aibridge/config" aibcontext "github.com/coder/coder/v2/aibridge/context" "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/aibridge/intercept/eventstream" @@ -36,7 +35,7 @@ func NewBlockingInterceptor( reqPayload RequestPayload, cfg intercept.Config, cred intercept.Credential, - bedrockCfg *aibconfig.AWSBedrock, + bedrock *BedrockRuntime, clientHeaders http.Header, tracer trace.Tracer, ) *BlockingInterception { @@ -45,7 +44,7 @@ func NewBlockingInterceptor( reqPayload: reqPayload, cfg: cfg, cred: cred, - bedrockCfg: bedrockCfg, + bedrock: bedrock, clientHeaders: clientHeaders, tracer: tracer, }} diff --git a/aibridge/intercept/messages/streaming.go b/aibridge/intercept/messages/streaming.go index 442c4291b0..2369c17654 100644 --- a/aibridge/intercept/messages/streaming.go +++ b/aibridge/intercept/messages/streaming.go @@ -21,7 +21,6 @@ import ( "golang.org/x/xerrors" "cdr.dev/slog/v3" - aibconfig "github.com/coder/coder/v2/aibridge/config" aibcontext "github.com/coder/coder/v2/aibridge/context" "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/aibridge/intercept/eventstream" @@ -41,7 +40,7 @@ func NewStreamingInterceptor( reqPayload RequestPayload, cfg intercept.Config, cred intercept.Credential, - bedrockCfg *aibconfig.AWSBedrock, + bedrock *BedrockRuntime, clientHeaders http.Header, tracer trace.Tracer, ) *StreamingInterception { @@ -50,7 +49,7 @@ func NewStreamingInterceptor( reqPayload: reqPayload, cfg: cfg, cred: cred, - bedrockCfg: bedrockCfg, + bedrock: bedrock, clientHeaders: clientHeaders, tracer: tracer, }} diff --git a/aibridge/internal/integrationtest/apidump_internal_test.go b/aibridge/internal/integrationtest/apidump_internal_test.go index 42811cb362..b48af4c3e7 100644 --- a/aibridge/internal/integrationtest/apidump_internal_test.go +++ b/aibridge/internal/integrationtest/apidump_internal_test.go @@ -15,6 +15,7 @@ import ( "github.com/stretchr/testify/require" "github.com/coder/coder/v2/aibridge" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/config" "github.com/coder/coder/v2/aibridge/fixtures" "github.com/coder/coder/v2/aibridge/intercept/apidump" @@ -39,7 +40,7 @@ func TestAPIDump(t *testing.T) { name: "anthropic", fixture: fixtures.AntSimple, providerFunc: func(addr, dumpDir string) aibridge.Provider { - return provider.NewAnthropic(anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) + return aibridgetest.NewAnthropicProvider(t, anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) }, path: pathAnthropicMessages, expectProviderDir: config.ProviderAnthropic, @@ -219,7 +220,7 @@ func TestAPIDumpPassthrough(t *testing.T) { { name: "anthropic", providerFunc: func(addr string, dumpDir string) aibridge.Provider { - return provider.NewAnthropic(anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) + return aibridgetest.NewAnthropicProvider(t, anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) }, requestPath: "/anthropic/v1/models", expectDumpName: "-v1-models-", diff --git a/aibridge/internal/integrationtest/bridge_internal_test.go b/aibridge/internal/integrationtest/bridge_internal_test.go index ef226db2b9..024652928e 100644 --- a/aibridge/internal/integrationtest/bridge_internal_test.go +++ b/aibridge/internal/integrationtest/bridge_internal_test.go @@ -32,6 +32,7 @@ import ( "golang.org/x/xerrors" "github.com/coder/coder/v2/aibridge" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/config" "github.com/coder/coder/v2/aibridge/fixtures" "github.com/coder/coder/v2/aibridge/intercept" @@ -316,19 +317,8 @@ func TestAWSBedrockIntegration(t *testing.T) { SmallFastModel: "test-haiku", } - bridgeServer := newBridgeTestServer(ctx, t, "http://unused", - withCustomProvider(provider.NewAnthropic(anthropicCfg("http://unused", apiKey), bedrockCfg)), - ) - - resp, err := bridgeServer.makeRequest(t, http.MethodPost, pathAnthropicMessages, fixtures.Request(t, fixtures.AntSingleBuiltinTool)) - require.NoError(t, err) - defer resp.Body.Close() - - require.Equal(t, http.StatusInternalServerError, resp.StatusCode) - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - require.Contains(t, string(body), "create anthropic client") - require.Contains(t, string(body), "region or base url required") + _, err := provider.NewAnthropic(ctx, anthropicCfg("http://unused", apiKey), bedrockCfg) + require.ErrorContains(t, err, "region or base url required") }) t.Run("/v1/messages", func(t *testing.T) { @@ -353,7 +343,7 @@ func TestAWSBedrockIntegration(t *testing.T) { } bridgeServer := newBridgeTestServer(ctx, t, upstream.URL, - withCustomProvider(provider.NewAnthropic(anthropicCfg(upstream.URL, apiKey), bedrockCfg)), + withCustomProvider(aibridgetest.NewAnthropicProvider(t, anthropicCfg(upstream.URL, apiKey), bedrockCfg)), ) // Make API call to aibridge for Anthropic /v1/messages, which will be routed via AWS Bedrock. @@ -483,7 +473,7 @@ func TestAWSBedrockIntegration(t *testing.T) { } bridgeServer := newBridgeTestServer(ctx, t, upstream.URL, - withCustomProvider(provider.NewAnthropic(anthropicCfg(upstream.URL, apiKey), bCfg)), + withCustomProvider(aibridgetest.NewAnthropicProvider(t, anthropicCfg(upstream.URL, apiKey), bCfg)), ) reqBody, err := sjson.SetBytes(fix.Request(), "stream", streaming) @@ -637,7 +627,7 @@ func TestAWSBedrockIntegration(t *testing.T) { bCfg.Region = region bridgeServer := newBridgeTestServer(ctx, t, mockEgressProxy.URL, - withCustomProvider(provider.NewAnthropic(anthropicCfg(mockEgressProxy.URL, apiKey), bCfg)), + withCustomProvider(aibridgetest.NewAnthropicProvider(t, anthropicCfg(mockEgressProxy.URL, apiKey), bCfg)), ) // Sends a bridge request through a mock egress proxy that @@ -2292,7 +2282,7 @@ func TestActorHeaders(t *testing.T) { createProviderFn: func(url, key string, sendHeaders bool) aibridge.Provider { cfg := anthropicCfg(url, key) cfg.SendActorHeaders = sendHeaders - return provider.NewAnthropic(cfg, nil) + return aibridgetest.NewAnthropicProvider(t, cfg, nil) }, fixture: fixtures.AntSimple, streaming: true, @@ -2303,7 +2293,7 @@ func TestActorHeaders(t *testing.T) { createProviderFn: func(url, key string, sendHeaders bool) aibridge.Provider { cfg := anthropicCfg(url, key) cfg.SendActorHeaders = sendHeaders - return provider.NewAnthropic(cfg, nil) + return aibridgetest.NewAnthropicProvider(t, cfg, nil) }, fixture: fixtures.AntSimple, streaming: false, diff --git a/aibridge/internal/integrationtest/circuit_breaker_internal_test.go b/aibridge/internal/integrationtest/circuit_breaker_internal_test.go index e9909b5ae3..57f9b27df3 100644 --- a/aibridge/internal/integrationtest/circuit_breaker_internal_test.go +++ b/aibridge/internal/integrationtest/circuit_breaker_internal_test.go @@ -16,6 +16,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/config" "github.com/coder/coder/v2/aibridge/internal/testutil" "github.com/coder/coder/v2/aibridge/metrics" @@ -70,7 +71,7 @@ func TestCircuitBreaker_FullRecoveryCycle(t *testing.T) { }, path: pathAnthropicMessages, createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { - return provider.NewAnthropic(config.Anthropic{ + return aibridgetest.NewAnthropicProvider(t, config.Anthropic{ BaseURL: baseURL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, @@ -237,7 +238,7 @@ func TestCircuitBreaker_HalfOpenFailure(t *testing.T) { }, path: pathAnthropicMessages, createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { - return provider.NewAnthropic(config.Anthropic{ + return aibridgetest.NewAnthropicProvider(t, config.Anthropic{ BaseURL: baseURL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, @@ -374,7 +375,7 @@ func TestCircuitBreaker_HalfOpenMaxRequests(t *testing.T) { }, path: pathAnthropicMessages, createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { - return provider.NewAnthropic(config.Anthropic{ + return aibridgetest.NewAnthropicProvider(t, config.Anthropic{ BaseURL: baseURL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, @@ -555,7 +556,7 @@ func TestCircuitBreaker_PerModelIsolation(t *testing.T) { } ctx := t.Context() bridgeServer := newBridgeTestServer(ctx, t, mockUpstream.URL, - withCustomProvider(provider.NewAnthropic(config.Anthropic{ + withCustomProvider(aibridgetest.NewAnthropicProvider(t, config.Anthropic{ BaseURL: mockUpstream.URL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, diff --git a/aibridge/internal/integrationtest/keypool_failover_internal_test.go b/aibridge/internal/integrationtest/keypool_failover_internal_test.go index af92056c87..cb6d3c24ef 100644 --- a/aibridge/internal/integrationtest/keypool_failover_internal_test.go +++ b/aibridge/internal/integrationtest/keypool_failover_internal_test.go @@ -10,6 +10,7 @@ import ( "github.com/tidwall/sjson" "github.com/coder/coder/v2/aibridge" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/config" "github.com/coder/coder/v2/aibridge/fixtures" "github.com/coder/coder/v2/aibridge/internal/testutil" @@ -158,7 +159,7 @@ func TestAnthropic_KeyFailover(t *testing.T) { ) bridgeServer := newBridgeTestServer(t.Context(), t, upstream.URL, - withCustomProvider(provider.NewAnthropic(config.Anthropic{ + withCustomProvider(aibridgetest.NewAnthropicProvider(t, config.Anthropic{ BaseURL: upstream.URL, KeyPool: pool, }, nil)), @@ -232,7 +233,7 @@ func TestKeyPool_StateSharing(t *testing.T) { name: "anthropic", providerName: config.ProviderAnthropic, newProvider: func(baseURL string, pool *keypool.Pool) aibridge.Provider { - return provider.NewAnthropic(config.Anthropic{BaseURL: baseURL, KeyPool: pool}, nil) + return aibridgetest.NewAnthropicProvider(t, config.Anthropic{BaseURL: baseURL, KeyPool: pool}, nil) }, upstreamResponses: []testutil.UpstreamResponse{ testutil.NewErrorResponse(http.StatusTooManyRequests, "60"), diff --git a/aibridge/internal/integrationtest/setupbridge.go b/aibridge/internal/integrationtest/setupbridge.go index e63f554a00..d2e8c0929b 100644 --- a/aibridge/internal/integrationtest/setupbridge.go +++ b/aibridge/internal/integrationtest/setupbridge.go @@ -16,6 +16,7 @@ import ( "cdr.dev/slog/v3" "github.com/coder/coder/v2/aibridge" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/config" aibcontext "github.com/coder/coder/v2/aibridge/context" "github.com/coder/coder/v2/aibridge/fixtures" @@ -45,7 +46,7 @@ const ( var defaultTracer = otel.Tracer("integrationtest") type bridgeConfig struct { - providerBuilders []func(upstreamURL string) aibridge.Provider + providerBuilders []func(t *testing.T, upstreamURL string) aibridge.Provider metrics *metrics.Metrics tracer trace.Tracer mcpProxy mcp.ServerProxier @@ -87,8 +88,8 @@ type bridgeOption func(*bridgeConfig) // When any provider option is used, the default "all providers" set is not created. func withProvider(providerType string) bridgeOption { return func(c *bridgeConfig) { - c.providerBuilders = append(c.providerBuilders, func(addr string) aibridge.Provider { - return newDefaultProvider(providerType, addr) + c.providerBuilders = append(c.providerBuilders, func(t *testing.T, addr string) aibridge.Provider { + return newDefaultProvider(t, providerType, addr) }) } } @@ -98,7 +99,7 @@ func withProvider(providerType string) bridgeOption { // When any provider option is used, the default "all providers" set is not created. func withCustomProvider(p aibridge.Provider) bridgeOption { return func(c *bridgeConfig) { - c.providerBuilders = append(c.providerBuilders, func(string) aibridge.Provider { + c.providerBuilders = append(c.providerBuilders, func(*testing.T, string) aibridge.Provider { return p }) } @@ -158,12 +159,12 @@ func newBridgeTestServer( var providers []aibridge.Provider if len(cfg.providerBuilders) > 0 { for _, b := range cfg.providerBuilders { - providers = append(providers, b(upstreamURL)) + providers = append(providers, b(t, upstreamURL)) } } else { providers = []aibridge.Provider{ - newDefaultProvider(config.ProviderAnthropic, upstreamURL), - newDefaultProvider(config.ProviderOpenAI, upstreamURL), + newDefaultProvider(t, config.ProviderAnthropic, upstreamURL), + newDefaultProvider(t, config.ProviderOpenAI, upstreamURL), } } @@ -250,14 +251,14 @@ func setupInjectedToolTest( } // newDefaultProvider creates a Provider with default test configuration. -func newDefaultProvider(providerType string, addr string) aibridge.Provider { +func newDefaultProvider(t *testing.T, providerType string, addr string) aibridge.Provider { switch providerType { case config.ProviderAnthropic: - return provider.NewAnthropic(anthropicCfg(addr, apiKey), nil) + return aibridgetest.NewAnthropicProvider(t, anthropicCfg(addr, apiKey), nil) case config.ProviderOpenAI: return provider.NewOpenAI(openAICfg(addr, apiKey)) case providerBedrock: - return provider.NewAnthropic(anthropicCfg(addr, apiKey), bedrockCfg(addr)) + return aibridgetest.NewAnthropicProvider(t, anthropicCfg(addr, apiKey), bedrockCfg(addr)) default: panic("unknown provider type: " + providerType) } diff --git a/aibridge/passthrough_internal_test.go b/aibridge/passthrough_internal_test.go index 5a0aeb3cb6..76ca6c1731 100644 --- a/aibridge/passthrough_internal_test.go +++ b/aibridge/passthrough_internal_test.go @@ -1,6 +1,7 @@ package aibridge import ( + "context" "crypto/tls" "maps" "net" @@ -316,10 +317,12 @@ func TestPassthrough_KeyFailover(t *testing.T) { r.Header.Set("X-Api-Key", key) }, newProvider: func(baseURL string, pool *keypool.Pool) provider.Provider { - return provider.NewAnthropic(config.Anthropic{ + p, err := provider.NewAnthropic(context.Background(), config.Anthropic{ BaseURL: baseURL, KeyPool: pool, }, nil) + require.NoError(t, err) + return p }, }, { diff --git a/aibridge/provider/anthropic.go b/aibridge/provider/anthropic.go index 2b04bdf8b1..01cc587ecf 100644 --- a/aibridge/provider/anthropic.go +++ b/aibridge/provider/anthropic.go @@ -1,6 +1,7 @@ package provider import ( + "context" "fmt" "io" "net/http" @@ -25,8 +26,9 @@ var _ Provider = &Anthropic{} // Anthropic allows for interactions with the Anthropic API. type Anthropic struct { - cfg config.Anthropic - bedrockCfg *config.AWSBedrock + cfg config.Anthropic + // bedrock is nil for non-Bedrock providers. + bedrock *messages.BedrockRuntime } const routeMessages = "/v1/messages" // https://docs.anthropic.com/en/api/messages @@ -43,7 +45,7 @@ var anthropicIsFailure = func(statusCode int) bool { return circuitbreaker.DefaultIsFailure(statusCode) } -func NewAnthropic(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) *Anthropic { +func NewAnthropic(ctx context.Context, cfg config.Anthropic, bedrockCfg *config.AWSBedrock) (*Anthropic, error) { if cfg.Name == "" { cfg.Name = config.ProviderAnthropic } @@ -55,10 +57,23 @@ func NewAnthropic(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) *Anthropi cfg.CircuitBreaker.OpenErrorResponse = anthropicOpenErrorResponse } - return &Anthropic{ - cfg: cfg, - bedrockCfg: bedrockCfg, + // Resolve the AWS credentials provider once and bundle it with the config. + // This performs no network call (the base identity and any AssumeRole + // resolve lazily on first retrieval); it only wires up the provider chain, + // so it is cheap to run at construction. + var bedrock *messages.BedrockRuntime + if bedrockCfg != nil { + creds, region, err := buildBedrockCredentials(ctx, *bedrockCfg) + if err != nil { + return nil, xerrors.Errorf("build bedrock credentials: %w", err) + } + bedrock = &messages.BedrockRuntime{Cfg: *bedrockCfg, Creds: creds, ResolvedRegion: region} } + + return &Anthropic{ + cfg: cfg, + bedrock: bedrock, + }, nil } func (*Anthropic) Type() string { @@ -123,9 +138,9 @@ func (p *Anthropic) CreateInterceptor(_ http.ResponseWriter, r *http.Request, tr var interceptor intercept.Interceptor if reqPayload.Stream() { - interceptor = messages.NewStreamingInterceptor(id, reqPayload, cfg, cred, p.bedrockCfg, r.Header, tracer) + interceptor = messages.NewStreamingInterceptor(id, reqPayload, cfg, cred, p.bedrock, r.Header, tracer) } else { - interceptor = messages.NewBlockingInterceptor(id, reqPayload, cfg, cred, p.bedrockCfg, r.Header, tracer) + interceptor = messages.NewBlockingInterceptor(id, reqPayload, cfg, cred, p.bedrock, r.Header, tracer) } span.SetAttributes(interceptor.TraceAttributes(r)...) return interceptor, nil @@ -153,8 +168,8 @@ func (p *Anthropic) resolveCredential(r *http.Request) (intercept.Credential, er if p.cfg.KeyPool != nil { return &intercept.CentralizedPool{Pool: p.cfg.KeyPool, Header: p.AuthHeader()}, nil } - if p.bedrockCfg != nil { - return intercept.Bedrock{AccessKey: p.bedrockCfg.AccessKey}, nil + if p.bedrock != nil { + return intercept.Bedrock{AccessKey: p.bedrock.Cfg.AccessKey}, nil } return nil, ErrNoCredential } diff --git a/aibridge/provider/anthropic_internal_test.go b/aibridge/provider/anthropic_internal_test.go index 679469fdc6..9df6f03718 100644 --- a/aibridge/provider/anthropic_internal_test.go +++ b/aibridge/provider/anthropic_internal_test.go @@ -2,6 +2,7 @@ package provider import ( "bytes" + "context" "net/http" "net/http/httptest" "testing" @@ -18,6 +19,16 @@ import ( "github.com/coder/quartz" ) +// newTestAnthropic is local (not aibridgetest.NewAnthropicProvider) because these +// white-box tests need the concrete *Anthropic, and importing aibridgetest here +// would create an import cycle. +func newTestAnthropic(t testing.TB, cfg config.Anthropic, bedrockCfg *config.AWSBedrock) *Anthropic { + t.Helper() + p, err := NewAnthropic(context.Background(), cfg, bedrockCfg) + require.NoError(t, err) + return p +} + func TestAnthropic_TypeAndName(t *testing.T) { t.Parallel() @@ -45,7 +56,7 @@ func TestAnthropic_TypeAndName(t *testing.T) { t.Run(tc.name, func(t *testing.T) { t.Parallel() - p := NewAnthropic(tc.cfg, nil) + p := newTestAnthropic(t, tc.cfg, nil) assert.Equal(t, tc.expectType, p.Type()) assert.Equal(t, tc.expectName, p.Name()) }) @@ -81,7 +92,7 @@ func TestNewAnthropic_KeyResolution(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { t.Parallel() - p := NewAnthropic(tc.cfg, nil) + p := newTestAnthropic(t, tc.cfg, nil) if tc.expectedKeys == nil { assert.Nil(t, p.cfg.KeyPool, "expected no KeyPool") @@ -106,7 +117,7 @@ func TestNewAnthropic_KeyResolution(t *testing.T) { func TestAnthropic_CreateInterceptor(t *testing.T) { t.Parallel() - provider := NewAnthropic(config.Anthropic{KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key")}, nil) + provider := newTestAnthropic(t, config.Anthropic{KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key")}, nil) t.Run("Messages_NonStreamingRequest_BlockingInterceptor", func(t *testing.T) { t.Parallel() @@ -164,7 +175,7 @@ func TestAnthropic_CreateInterceptor(t *testing.T) { })) t.Cleanup(mockUpstream.Close) - provider := NewAnthropic(config.Anthropic{ + provider := newTestAnthropic(t, config.Anthropic{ BaseURL: mockUpstream.URL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), }, nil) @@ -286,7 +297,7 @@ func TestAnthropic_CreateInterceptor_Credential(t *testing.T) { bedrock: true, setHeaders: map[string]string{}, wantCredentialKind: intercept.CredentialKindCentralized, - wantCredentialHint: "", + wantCredentialHint: "", }, { // Bedrock static mode: the hint masks the access key ID. @@ -331,7 +342,7 @@ func TestAnthropic_CreateInterceptor_Credential(t *testing.T) { bedrock.AccessKeySecret = "wJalrXUtnFEMI-secret-value" } } - provider := NewAnthropic(acfg, bedrock) + provider := newTestAnthropic(t, acfg, bedrock) body := `{"model": "claude-opus-4-5", "max_tokens": 1024, "messages": [{"role": "user", "content": "hello"}], "stream": false}` req := httptest.NewRequest(http.MethodPost, routeMessages, bytes.NewBufferString(body)) @@ -375,7 +386,7 @@ func TestAnthropic_KeyFailoverConfig(t *testing.T) { pool, err := keypool.New(config.ProviderAnthropic, []string{"k0", "k1"}, quartz.NewMock(t), nil) require.NoError(t, err) - p := NewAnthropic(config.Anthropic{KeyPool: pool}, nil) + p := newTestAnthropic(t, config.Anthropic{KeyPool: pool}, nil) cfg := p.KeyFailoverConfig(slog.Make()) diff --git a/aibridge/provider/bedrock.go b/aibridge/provider/bedrock.go new file mode 100644 index 0000000000..93bbf7038a --- /dev/null +++ b/aibridge/provider/bedrock.go @@ -0,0 +1,89 @@ +package provider + +import ( + "context" + + "github.com/aws/aws-sdk-go-v2/aws" + awsconfig "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/credentials" + "github.com/aws/aws-sdk-go-v2/credentials/stscreds" + "github.com/aws/aws-sdk-go-v2/service/sts" + "golang.org/x/xerrors" + + "github.com/coder/coder/v2/aibridge/config" +) + +// bedrockSessionName is the STS role session name attached to AssumeRole calls. +// A stable value keeps them identifiable in CloudTrail. +const bedrockSessionName = "coder-aigateway" + +// buildBedrockCredentials resolves the base identity and, when a role ARN +// is configured, assumes that role via STS. The base identity is either +// static keys or the AWS SDK default credential chain, which covers IRSA, +// EKS Pod Identity, EC2 Instance Profile, and more. +// +// The result is wrapped in aws.NewCredentialsCache, which caches and rotates +// the resolved temporary credentials. buildBedrockCredentials should be called +// once when the Bedrock provider is constructed, and the returned Credential +// Provider should be shared across all LLM requests to the Bedrock Provider, +// so per-request credential retrieval is served from this cache rather than +// re-resolving (and re-assuming) on every request. No network call is made here: +// the base identity and any AssumeRole are resolved lazily on first retrieval. +func buildBedrockCredentials(ctx context.Context, cfg config.AWSBedrock) (aws.CredentialsProvider, string, error) { + if cfg.Region == "" && cfg.BaseURL == "" { + return nil, "", xerrors.New("region or base url required") + } + + var loadOpts []func(*awsconfig.LoadOptions) error + if cfg.Region != "" { + loadOpts = append(loadOpts, awsconfig.WithRegion(cfg.Region)) + } + + // Use static credentials when explicitly provided, otherwise fall back to + // the SDK default credential chain. + switch { + // Both set: use static credentials directly. + case cfg.AccessKey != "" && cfg.AccessKeySecret != "": + loadOpts = append(loadOpts, awsconfig.WithCredentialsProvider( + credentials.NewStaticCredentialsProvider( + cfg.AccessKey, + cfg.AccessKeySecret, + "", + ), + )) + // Only one set: misconfiguration. + case cfg.AccessKey != "" || cfg.AccessKeySecret != "": + return nil, "", xerrors.New("both access key and access key secret must be provided together") + // Neither set: SDK default credential chain resolves the base identity. + default: + } + + base, err := awsconfig.LoadDefaultConfig(ctx, loadOpts...) + if err != nil { + return nil, "", xerrors.Errorf("failed to load AWS Bedrock config: %w", err) + } + + // Assuming a role calls STS, which needs a region to resolve its endpoint. + // The region may come from the config or the AWS environment; if neither + // supplies one, fail here. + if cfg.RoleARN != "" && base.Region == "" { + return nil, "", xerrors.New("region is required to assume a role, but was not specified") + } + + // The base identity signs Bedrock requests directly unless a target role is + // configured, in which case it signs the AssumeRole call and the resulting + // temporary credentials sign Bedrock requests. The default credential chain + // is already cache-wrapped, so only the AssumeRoleProvider is wrapped with a + // cache to avoid re-assuming the role on every request. + credsProvider := base.Credentials + if cfg.RoleARN != "" { + credsProvider = stscreds.NewAssumeRoleProvider(sts.NewFromConfig(base), cfg.RoleARN, func(o *stscreds.AssumeRoleOptions) { + o.RoleSessionName = bedrockSessionName + }) + credsProvider = aws.NewCredentialsCache(credsProvider) + } + + // base.Region is the region the SDK resolved (explicit config, AWS_REGION / + // AWS_DEFAULT_REGION, shared config, or IMDS). + return credsProvider, base.Region, nil +} diff --git a/aibridge/provider/bedrock_internal_test.go b/aibridge/provider/bedrock_internal_test.go new file mode 100644 index 0000000000..0ecae6c780 --- /dev/null +++ b/aibridge/provider/bedrock_internal_test.go @@ -0,0 +1,378 @@ +package provider + +import ( + "context" + "net/http" + "net/http/httptest" + "sync/atomic" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/aibridge/config" +) + +// TestBuildBedrockCredentialsValidation covers the input validation that does +// not require resolving credentials. +func TestBuildBedrockCredentialsValidation(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + cfg config.AWSBedrock + errorMsg string + }{ + { + name: "missing region and base url", + cfg: config.AWSBedrock{}, + errorMsg: "region or base url required", + }, + { + name: "missing access key", + cfg: config.AWSBedrock{ + Region: "us-east-1", + AccessKeySecret: "test-secret", + }, + errorMsg: "both access key and access key secret must be provided together", + }, + { + name: "missing access key secret", + cfg: config.AWSBedrock{ + Region: "us-east-1", + AccessKey: "test-key", + }, + errorMsg: "both access key and access key secret must be provided together", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + _, _, err := buildBedrockCredentials(context.Background(), tt.cfg) + require.Error(t, err) + require.Contains(t, err.Error(), tt.errorMsg) + }) + } +} + +// TestBuildBedrockCredentialsStatic resolves static credentials. +func TestBuildBedrockCredentialsStatic(t *testing.T) { + t.Parallel() + + creds, _, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + Region: "us-east-1", + AccessKey: "test-key", + AccessKeySecret: "test-secret", + }) + require.NoError(t, err) + + got, err := creds.Retrieve(context.Background()) + require.NoError(t, err) + require.Equal(t, "test-key", got.AccessKeyID) + require.Equal(t, "test-secret", got.SecretAccessKey) +} + +// TestBuildBedrockCredentialsDefaultChain covers resolution via the AWS SDK +// default credential chain. +// NOTE: no t.Parallel() because the subtests use t.Setenv. +func TestBuildBedrockCredentialsDefaultChain(t *testing.T) { + tests := []struct { + name string + envVars map[string]string + expectError bool + wantKey string + wantSecret string + wantToken string + }{ + { + name: "credentials via env", + envVars: map[string]string{ + "AWS_ACCESS_KEY_ID": "test-key", + "AWS_SECRET_ACCESS_KEY": "test-secret", + }, + wantKey: "test-key", + wantSecret: "test-secret", + }, + { + name: "credentials with session token via env", + envVars: map[string]string{ + "AWS_ACCESS_KEY_ID": "test-key", + "AWS_SECRET_ACCESS_KEY": "test-secret", + "AWS_SESSION_TOKEN": "test-session-token", + }, + wantKey: "test-key", + wantSecret: "test-secret", + wantToken: "test-session-token", + }, + { + name: "error when no credential source is configured", + envVars: map[string]string{ + "AWS_ACCESS_KEY_ID": "", + "AWS_SECRET_ACCESS_KEY": "", + "AWS_SESSION_TOKEN": "", + "AWS_PROFILE": "", + "AWS_SHARED_CREDENTIALS_FILE": "/dev/null", + "AWS_CONFIG_FILE": "/dev/null", + "AWS_WEB_IDENTITY_TOKEN_FILE": "", + "AWS_ROLE_ARN": "", + "AWS_ROLE_SESSION_NAME": "", + "AWS_CONTAINER_CREDENTIALS_RELATIVE_URI": "", + "AWS_CONTAINER_CREDENTIALS_FULL_URI": "", + "AWS_CONTAINER_AUTHORIZATION_TOKEN": "", + "AWS_EC2_METADATA_DISABLED": "true", + }, + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + for key, val := range tt.envVars { + t.Setenv(key, val) + } + + // buildBedrockCredentials only wires up the provider chain; it + // does not resolve credentials, so it succeeds regardless of + // credential availability. Resolution failures surface on Retrieve. + creds, _, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + Region: "us-east-1", + }) + require.NoError(t, err) + require.NotNil(t, creds) + + got, err := creds.Retrieve(context.Background()) + if tt.expectError { + require.Error(t, err) + return + } + require.NoError(t, err) + require.Equal(t, tt.wantKey, got.AccessKeyID) + require.Equal(t, tt.wantSecret, got.SecretAccessKey) + require.Equal(t, tt.wantToken, got.SessionToken) + }) + } +} + +// TestBuildBedrockCredentialsAssumeRole drives the STS AssumeRole path against a +// mock endpoint, asserting that the configured role ARN and the stable session +// name are sent and that the returned temporary credentials are used. +// NOTE: no t.Parallel() because it uses t.Setenv. +func TestBuildBedrockCredentialsAssumeRole(t *testing.T) { + var gotRoleARN, gotSessionName string + // Mock the AWS STS AssumeRole API. + // https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRole.html + sts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.NoError(t, r.ParseForm()) + gotRoleARN = r.Form.Get("RoleArn") + gotSessionName = r.Form.Get("RoleSessionName") + + w.Header().Set("Content-Type", "text/xml") + _, _ = w.Write([]byte(` + + + ASIAASSUMED + assumed-secret + assumed-token + 2999-01-01T00:00:00Z + + + arn:aws:sts::123456789012:assumed-role/target/coder + AROAEXAMPLE:coder + + +`)) + })) + defer sts.Close() + + // Point the STS client at the mock and provide static base credentials so + // the base identity resolves without additional network calls. + t.Setenv("AWS_ENDPOINT_URL_STS", sts.URL) + t.Setenv("AWS_ACCESS_KEY_ID", "base-key") + t.Setenv("AWS_SECRET_ACCESS_KEY", "base-secret") + + creds, _, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + Region: "us-east-1", + RoleARN: "arn:aws:iam::123456789012:role/target", + }) + require.NoError(t, err) + + got, err := creds.Retrieve(context.Background()) + require.NoError(t, err) + require.Equal(t, "ASIAASSUMED", got.AccessKeyID) + require.Equal(t, "assumed-secret", got.SecretAccessKey) + require.Equal(t, "assumed-token", got.SessionToken) + + require.Equal(t, "arn:aws:iam::123456789012:role/target", gotRoleARN) + require.Equal(t, bedrockSessionName, gotSessionName) +} + +// TestBuildBedrockCredentialsAssumeRoleError verifies that when STS rejects the +// AssumeRole call (e.g. a trust-policy or IAM denial), the failure surfaces to +// the caller on Retrieve with enough detail to diagnose it, rather than being +// swallowed. The base identity resolved fine; only the role assumption failed. +// NOTE: no t.Parallel() because it uses t.Setenv. +func TestBuildBedrockCredentialsAssumeRoleError(t *testing.T) { + // Mock the AWS STS AssumeRole API. + // https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRole.html + sts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/xml") + w.WriteHeader(http.StatusForbidden) + _, _ = w.Write([]byte(` + + Sender + AccessDenied + User arn:aws:iam::123456789012:user/base is not authorized to perform sts:AssumeRole on arn:aws:iam::123456789012:role/target + +`)) + })) + defer sts.Close() + + t.Setenv("AWS_ENDPOINT_URL_STS", sts.URL) + t.Setenv("AWS_ACCESS_KEY_ID", "base-key") + t.Setenv("AWS_SECRET_ACCESS_KEY", "base-secret") + + creds, _, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + Region: "us-east-1", + RoleARN: "arn:aws:iam::123456789012:role/target", + }) + require.NoError(t, err) // Build is lazy; the STS call happens on Retrieve. + + _, err = creds.Retrieve(context.Background()) + require.Error(t, err) + // The error must carry the STS operation and failure code so operators can + // tell this is an AssumeRole authorization problem, not missing credentials. + require.ErrorContains(t, err, "AssumeRole") + require.ErrorContains(t, err, "AccessDenied") +} + +// TestBuildBedrockCredentialsAssumeRoleCaches verifies the AssumeRole result is +// cached: many credential retrievals, one per LLM request, trigger a single STS +// AssumeRole call rather than re-assuming the role on every request. +// NOTE: no t.Parallel() because it uses t.Setenv. +func TestBuildBedrockCredentialsAssumeRoleCaches(t *testing.T) { + var stsCalls atomic.Int64 + // Mock the AWS STS AssumeRole API. + // https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRole.html + sts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + stsCalls.Add(1) + w.Header().Set("Content-Type", "text/xml") + // A far-future expiration keeps the cached credentials valid, so the + // cache serves every retrieval after the first without re-assuming. + _, _ = w.Write([]byte(` + + + ASIAASSUMED + assumed-secret + assumed-token + 2999-01-01T00:00:00Z + + +`)) + })) + defer sts.Close() + + t.Setenv("AWS_ENDPOINT_URL_STS", sts.URL) + t.Setenv("AWS_ACCESS_KEY_ID", "base-key") + t.Setenv("AWS_SECRET_ACCESS_KEY", "base-secret") + + creds, _, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + Region: "us-east-1", + RoleARN: "arn:aws:iam::123456789012:role/target", + }) + require.NoError(t, err) + + // Each retrieval stands in for an LLM request resolving credentials from the + // shared provider. Only the first should reach STS. + for range 5 { + got, err := creds.Retrieve(context.Background()) + require.NoError(t, err) + require.Equal(t, "ASIAASSUMED", got.AccessKeyID) + } + + require.Equal(t, int64(1), stsCalls.Load(), + "AssumeRole should be called once, then served from the credentials cache") +} + +// TestBuildBedrockCredentialsAssumeRoleRefreshesOnExpiry verifies that once the +// assumed credentials expire, the next retrieval re-assumes the role rather than +// serving stale credentials from the cache. +// NOTE: no t.Parallel() because it uses t.Setenv. +func TestBuildBedrockCredentialsAssumeRoleRefreshesOnExpiry(t *testing.T) { + var stsCalls atomic.Int64 + // Mock the AWS STS AssumeRole API. + // https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRole.html + sts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + stsCalls.Add(1) + w.Header().Set("Content-Type", "text/xml") + // An expiration in the past makes the returned credentials immediately + // stale, so the cache cannot reuse them and must re-assume on the next + // retrieval. + _, _ = w.Write([]byte(` + + + ASIAASSUMED + assumed-secret + assumed-token + 2000-01-01T00:00:00Z + + +`)) + })) + defer sts.Close() + + t.Setenv("AWS_ENDPOINT_URL_STS", sts.URL) + t.Setenv("AWS_ACCESS_KEY_ID", "base-key") + t.Setenv("AWS_SECRET_ACCESS_KEY", "base-secret") + + creds, _, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + Region: "us-east-1", + RoleARN: "arn:aws:iam::123456789012:role/target", + }) + require.NoError(t, err) + + _, err = creds.Retrieve(context.Background()) + require.NoError(t, err) + _, err = creds.Retrieve(context.Background()) + require.NoError(t, err) + + require.Equal(t, int64(2), stsCalls.Load(), + "expired credentials should trigger a fresh AssumeRole on the next retrieval") +} + +// TestBuildBedrockCredentialsAssumeRoleRequiresRegion verifies that configuring +// a role without a resolvable region fails at construction. STS needs a region +// to resolve its endpoint. +// NOTE: no t.Parallel() because it uses t.Setenv. +func TestBuildBedrockCredentialsAssumeRoleRequiresRegion(t *testing.T) { + // Ensure no region resolves from the environment, shared config, or IMDS, + // so base.Region ends up empty. + t.Setenv("AWS_REGION", "") + t.Setenv("AWS_DEFAULT_REGION", "") + t.Setenv("AWS_PROFILE", "") + t.Setenv("AWS_CONFIG_FILE", "/dev/null") + t.Setenv("AWS_SHARED_CREDENTIALS_FILE", "/dev/null") + t.Setenv("AWS_EC2_METADATA_DISABLED", "true") + + _, _, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + BaseURL: "https://bedrock-runtime.example.com", + RoleARN: "arn:aws:iam::123456789012:role/target", + }) + require.ErrorContains(t, err, "region is required to assume a role") +} + +// TestBuildBedrockCredentialsAssumeRoleRegionFromEnv verifies that a role +// configured without an explicit region resolves it from the AWS environment +// (AWS_REGION here). +// NOTE: no t.Parallel() because it uses t.Setenv. +func TestBuildBedrockCredentialsAssumeRoleRegionFromEnv(t *testing.T) { + t.Setenv("AWS_REGION", "us-west-2") + + // BaseURL set with no explicit region: the region comes from AWS_REGION. + _, region, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + BaseURL: "https://bedrock-runtime.example.com", + RoleARN: "arn:aws:iam::123456789012:role/target", + }) + require.NoError(t, err) + require.Equal(t, "us-west-2", region) +} diff --git a/cli/aibridged.go b/cli/aibridged.go index e838c1db5a..9aa2ea5c27 100644 --- a/cli/aibridged.go +++ b/cli/aibridged.go @@ -175,7 +175,7 @@ func BuildProviders(ctx context.Context, db database.Store, cfg codersdk.AIBridg if row.Enabled { enabledCount++ } - prov, err := buildAIProviderFromRow(row, keysByProvider[row.ID], cfg, metrics) + prov, err := buildAIProviderFromRow(ctx, row, keysByProvider[row.ID], cfg, metrics) if err != nil { outcome.Status = aibridged.ProviderStatusError outcome.Err = err @@ -210,6 +210,7 @@ func BuildProviders(ctx context.Context, db database.Store, cfg codersdk.AIBridg // Disabled: true; settings decode, key loading, and credential checks // are skipped because the provider will never call upstream. func buildAIProviderFromRow( + ctx context.Context, row database.AIProvider, keys []database.AIProviderKey, cfg codersdk.AIBridgeConfig, @@ -284,14 +285,14 @@ func buildAIProviderFromRow( return nil, xerrors.Errorf("anthropic key pool: %w", err) } } - return aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{ + return aibridge.NewAnthropicProvider(ctx, aibridge.AnthropicConfig{ Name: row.Name, BaseURL: row.BaseUrl, KeyPool: pool, APIDumpDir: dumpDir, CircuitBreaker: cbCfg, SendActorHeaders: sendActorHeaders, - }, bedrock), nil + }, bedrock) case database.AIProviderTypeCopilot: // Copilot is always BYOK; the per-user token is supplied on each @@ -348,6 +349,7 @@ func bedrockConfigFromRow(row database.AIProvider, settings codersdk.AIProviderS AccessKeySecret: accessKeySecret, Model: bedrockSettings.Model, SmallFastModel: bedrockSettings.SmallFastModel, + RoleARN: bedrockSettings.RoleARN, } } diff --git a/cli/aibridged_internal_test.go b/cli/aibridged_internal_test.go index dee00f841f..d431b063b7 100644 --- a/cli/aibridged_internal_test.go +++ b/cli/aibridged_internal_test.go @@ -267,6 +267,7 @@ func TestBuildProviders(t *testing.T) { Name: "anthropic-bedrock", BaseUrl: "https://bedrock-runtime.us-west-2.amazonaws.com/", } + roleARN := "arn:aws:iam::123456789012:role/BedrockRole" settings := codersdk.AIProviderSettings{ Bedrock: &codersdk.AIProviderBedrockSettings{ Region: "us-west-2", @@ -274,6 +275,7 @@ func TestBuildProviders(t *testing.T) { AccessKeySecret: &secret, Model: model, SmallFastModel: smallModel, + RoleARN: roleARN, }, } got := bedrockConfigFromRow(row, settings) @@ -284,6 +286,7 @@ func TestBuildProviders(t *testing.T) { assert.Equal(t, secret, got.AccessKeySecret) assert.Equal(t, model, got.Model) assert.Equal(t, smallModel, got.SmallFastModel) + assert.Equal(t, roleARN, got.RoleARN) }) t.Run("BedrockSettingsEmpty", func(t *testing.T) { diff --git a/cli/server_aibridge_internal_test.go b/cli/server_aibridge_internal_test.go index 21711b0289..cb08ec530b 100644 --- a/cli/server_aibridge_internal_test.go +++ b/cli/server_aibridge_internal_test.go @@ -741,7 +741,7 @@ func TestBuildAIProviderFromRowSetsAPIDumpDir(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - provider, err := buildAIProviderFromRow(tt.row, nil, codersdk.AIBridgeConfig{ + provider, err := buildAIProviderFromRow(t.Context(), tt.row, nil, codersdk.AIBridgeConfig{ AllowBYOK: serpent.Bool(true), APIDumpDir: serpent.String(dumpDir), }, nil) @@ -755,7 +755,7 @@ func TestBuildAIProviderFromRowSetsAPIDumpDir(t *testing.T) { func TestBuildAIProviderFromRowBedrockWithoutSettings(t *testing.T) { t.Parallel() - _, err := buildAIProviderFromRow(database.AIProvider{ + _, err := buildAIProviderFromRow(t.Context(), database.AIProvider{ Enabled: true, Type: database.AIProviderTypeBedrock, Name: "bedrock-no-settings", diff --git a/coderd/ai_providers_test.go b/coderd/ai_providers_test.go index b9bfd283f1..4fcc1dd63b 100644 --- a/coderd/ai_providers_test.go +++ b/coderd/ai_providers_test.go @@ -1501,4 +1501,56 @@ func TestAIProviderSettingsMerge(t *testing.T) { require.NotNil(t, persisted.Bedrock.AccessKeySecret) require.Equal(t, "secret-new", *persisted.Bedrock.AccessKeySecret) }) + + t.Run("MigrateStaticToRole", func(t *testing.T) { + t.Parallel() + // An admin migrating from static AWS credentials to IAM role assumption + // clears the keys and sets a role ARN in a single PATCH. + client, db := coderdtest.NewWithDatabase(t, nil) + _ = coderdtest.CreateFirstUser(t, client) + ctx := testutil.Context(t, testutil.WaitLong) + + //nolint:gocritic // Owner role is the audience for this endpoint. + created, err := client.CreateAIProvider(ctx, codersdk.CreateAIProviderRequest{ + Type: codersdk.AIProviderTypeAnthropic, + Name: "merge-role", + Enabled: true, + BaseURL: "https://bedrock-runtime.us-east-1.amazonaws.com/", + Settings: codersdk.AIProviderSettings{ + Bedrock: &codersdk.AIProviderBedrockSettings{ + Region: "us-east-1", + AccessKey: ptr.Ref("AKIA-old"), //nolint:gosec // test fixture, not a real credential + AccessKeySecret: ptr.Ref("secret-old"), + }, + }, + }) + require.NoError(t, err) + + updated, err := client.UpdateAIProvider(ctx, created.Name, codersdk.UpdateAIProviderRequest{ + Settings: &codersdk.AIProviderSettings{ + Bedrock: &codersdk.AIProviderBedrockSettings{ + Region: "us-east-1", + AccessKey: ptr.Ref(""), + AccessKeySecret: ptr.Ref(""), + RoleARN: "arn:aws:iam::123456789012:role/target", + }, + }, + }) + require.NoError(t, err) + + require.NotNil(t, updated.Settings.Bedrock) + require.Equal(t, "arn:aws:iam::123456789012:role/target", updated.Settings.Bedrock.RoleARN) + + //nolint:gocritic // Test reads the row to verify write-only fields. + row, err := db.GetAIProviderByID(dbauthz.AsSystemRestricted(ctx), created.ID) + require.NoError(t, err) + persisted, err := db2sdk.AIProviderSettings(row.Settings) + require.NoError(t, err) + require.NotNil(t, persisted.Bedrock) + require.Equal(t, "arn:aws:iam::123456789012:role/target", persisted.Bedrock.RoleARN) + require.NotNil(t, persisted.Bedrock.AccessKey) + require.Equal(t, "", *persisted.Bedrock.AccessKey) + require.NotNil(t, persisted.Bedrock.AccessKeySecret) + require.Equal(t, "", *persisted.Bedrock.AccessKeySecret) + }) } diff --git a/coderd/aibridged/aibridged_test.go b/coderd/aibridged/aibridged_test.go index c0592e7390..e1fd056075 100644 --- a/coderd/aibridged/aibridged_test.go +++ b/coderd/aibridged/aibridged_test.go @@ -17,6 +17,7 @@ import ( "cdr.dev/slog/v3/sloggers/slogtest" "github.com/coder/coder/v2/aibridge" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/aibridge/keypool" agplaibridge "github.com/coder/coder/v2/coderd/aibridge" @@ -659,7 +660,7 @@ func TestServeHTTP_ActorHeaders(t *testing.T) { KeyPool: singleKeyPool(t, "openai", "test-key"), SendActorHeaders: true, }), - aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{ + aibridgetest.NewAnthropicProvider(t, aibridge.AnthropicConfig{ BaseURL: upstreamSrv.URL, KeyPool: singleKeyPool(t, "anthropic", "test-key"), SendActorHeaders: true, @@ -766,7 +767,7 @@ func TestRouting(t *testing.T) { providers := []aibridge.Provider{ aibridge.NewOpenAIProvider(aibridge.OpenAIConfig{BaseURL: openaiSrv.URL, KeyPool: singleKeyPool(t, "openai", "test-key")}), - aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{BaseURL: antSrv.URL, KeyPool: singleKeyPool(t, "anthropic", "test-key")}, nil), + aibridgetest.NewAnthropicProvider(t, aibridge.AnthropicConfig{BaseURL: antSrv.URL, KeyPool: singleKeyPool(t, "anthropic", "test-key")}, nil), } pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, providers, logger, nil, testTracer) require.NoError(t, err) diff --git a/codersdk/aiproviders.go b/codersdk/aiproviders.go index 7b513340bc..6cf8fb3359 100644 --- a/codersdk/aiproviders.go +++ b/codersdk/aiproviders.go @@ -11,6 +11,7 @@ import ( "strings" "time" + "github.com/aws/aws-sdk-go-v2/aws/arn" "github.com/google/uuid" "golang.org/x/xerrors" ) @@ -246,6 +247,9 @@ func (req CreateAIProviderRequest) Validate() []ValidationError { Detail: "type=bedrock does not accept api_keys", }) } + if req.Settings.Bedrock != nil { + validations = append(validations, validateAIProviderRoleARN(req.Settings.Bedrock.RoleARN)...) + } if req.Type == AIProviderTypeCopilot && len(req.APIKeys) > 0 { validations = append(validations, ValidationError{ Field: "api_keys", @@ -294,6 +298,9 @@ func (req UpdateAIProviderRequest) Validate() []ValidationError { if req.APIKeys != nil { validations = append(validations, validateAIProviderKeyMutations(*req.APIKeys)...) } + if req.Settings != nil && req.Settings.Bedrock != nil { + validations = append(validations, validateAIProviderRoleARN(req.Settings.Bedrock.RoleARN)...) + } return validations } @@ -316,6 +323,27 @@ func validateAIProviderName(name string) []ValidationError { return validations } +func validateAIProviderRoleARN(roleARN string) []ValidationError { + if roleARN == "" { + return nil + } + const exampleRoleARN = "arn:aws:iam::123456789012:role/BedrockRole" + invalid := func(detail string) []ValidationError { + return []ValidationError{{Field: "settings.role_arn", Detail: detail}} + } + parsed, err := arn.Parse(roleARN) + if err != nil { + return invalid(fmt.Sprintf("role_arn %q is not a valid ARN, e.g. %s", roleARN, exampleRoleARN)) + } + if parsed.Service != "iam" { + return invalid(fmt.Sprintf("role_arn must be an IAM ARN, but resolved to service %q, e.g. %s", parsed.Service, exampleRoleARN)) + } + if !strings.HasPrefix(parsed.Resource, "role/") { + return invalid(fmt.Sprintf("role_arn must reference an IAM role, but resolved to resource %q, e.g. %s", parsed.Resource, exampleRoleARN)) + } + return nil +} + func validateRequiredAIProviderBaseURL(raw string) []ValidationError { if raw == "" { return []ValidationError{{Field: "base_url", Detail: "base_url is required"}} diff --git a/codersdk/aiproviders_bedrock.go b/codersdk/aiproviders_bedrock.go index 88edcb0017..360b763333 100644 --- a/codersdk/aiproviders_bedrock.go +++ b/codersdk/aiproviders_bedrock.go @@ -30,6 +30,11 @@ type AIProviderBedrockSettings struct { // AccessKeySecret is the AWS secret access key paired with // AccessKey. Write-only. AccessKeySecret *string `json:"access_key_secret,omitempty"` + // RoleARN, when set, is the IAM role assumed via STS before calling + // Bedrock. The base identity (static keys or the AWS environment, e.g. + // IRSA / EKS Pod Identity / EC2 Instance Profile) signs the AssumeRole + // call, and the resulting temporary credentials sign Bedrock requests. + RoleARN string `json:"role_arn,omitempty"` } // IsConfigured reports whether any load-bearing Bedrock field is set, @@ -47,6 +52,9 @@ func (b AIProviderBedrockSettings) IsConfigured() bool { if b.Region != "" { return true } + if b.RoleARN != "" { + return true + } if b.AccessKey != nil && *b.AccessKey != "" { return true } diff --git a/codersdk/aiproviders_test.go b/codersdk/aiproviders_test.go index 97baad6535..a0ce28ec4e 100644 --- a/codersdk/aiproviders_test.go +++ b/codersdk/aiproviders_test.go @@ -159,3 +159,56 @@ func TestAIProviderSettings_Roundtrip(t *testing.T) { require.NoError(t, json.Unmarshal(encoded, &got)) require.Equal(t, orig, got) } + +func TestAIProviderRequest_ValidateRoleARN(t *testing.T) { + t.Parallel() + + cases := []struct { + name string + roleARN string + wantErr bool + }{ + {name: "empty is allowed", roleARN: "", wantErr: false}, + {name: "standard role arn", roleARN: "arn:aws:iam::743809215448:role/bedrock-role", wantErr: false}, + {name: "govcloud partition", roleARN: "arn:aws-us-gov:iam::123456789012:role/bedrock-role", wantErr: false}, + {name: "china partition", roleARN: "arn:aws-cn:iam::123456789012:role/bedrock-role", wantErr: false}, + {name: "role path", roleARN: "arn:aws:iam::123456789012:role/team/bedrock-role", wantErr: false}, + {name: "not an arn", roleARN: "bedrock-role", wantErr: true}, + {name: "wrong resource type", roleARN: "arn:aws:iam::123456789012:user/dave", wantErr: true}, + {name: "wrong service", roleARN: "arn:aws:s3:::my-bucket", wantErr: true}, + {name: "truncated arn", roleARN: "arn:aws:iam::123456789012", wantErr: true}, + } + + hasRoleARNError := func(vs []codersdk.ValidationError) bool { + for _, v := range vs { + if v.Field == "settings.role_arn" { + return true + } + } + return false + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + settings := codersdk.AIProviderSettings{ + Bedrock: &codersdk.AIProviderBedrockSettings{ + Region: "us-east-1", + RoleARN: tc.roleARN, + }, + } + + create := codersdk.CreateAIProviderRequest{ + Type: codersdk.AIProviderTypeBedrock, + Name: "bedrock", + BaseURL: "https://bedrock-runtime.us-east-1.amazonaws.com", + Settings: settings, + } + require.Equal(t, tc.wantErr, hasRoleARNError(create.Validate())) + + update := codersdk.UpdateAIProviderRequest{Settings: &settings} + require.Equal(t, tc.wantErr, hasRoleARNError(update.Validate())) + }) + } +} diff --git a/enterprise/aibridged_integration_test.go b/enterprise/aibridged_integration_test.go index 791c7323a5..f370d8a2a6 100644 --- a/enterprise/aibridged_integration_test.go +++ b/enterprise/aibridged_integration_test.go @@ -20,6 +20,7 @@ import ( "go.opentelemetry.io/otel/sdk/trace/tracetest" "github.com/coder/coder/v2/aibridge" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/config" "github.com/coder/coder/v2/aibridge/keypool" aibtracing "github.com/coder/coder/v2/aibridge/tracing" @@ -504,7 +505,7 @@ func TestIntegrationCircuitBreaker(t *testing.T) { KeyPool: singleKeyPool(t, config.ProviderOpenAI, "test-key"), CircuitBreaker: cbConfig, }), - aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{ + aibridgetest.NewAnthropicProvider(t, aibridge.AnthropicConfig{ BaseURL: mockAnthropic.URL, KeyPool: singleKeyPool(t, config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, diff --git a/go.mod b/go.mod index c3ea78e67a..ede6490189 100644 --- a/go.mod +++ b/go.mod @@ -311,7 +311,7 @@ require ( github.com/aws/aws-sdk-go-v2/service/ssm v1.67.4 // indirect github.com/aws/aws-sdk-go-v2/service/sso v1.31.3 // indirect github.com/aws/aws-sdk-go-v2/service/ssooidc v1.36.6 // indirect - github.com/aws/aws-sdk-go-v2/service/sts v1.43.3 // indirect + github.com/aws/aws-sdk-go-v2/service/sts v1.43.3 github.com/aymanbagabas/go-osc52/v2 v2.0.1 // indirect github.com/aymerick/douceur v0.2.0 // indirect github.com/beorn7/perks v1.0.1 // indirect diff --git a/site/src/api/typesGenerated.ts b/site/src/api/typesGenerated.ts index 9301a9eb79..6ecb349f45 100644 --- a/site/src/api/typesGenerated.ts +++ b/site/src/api/typesGenerated.ts @@ -322,6 +322,13 @@ export interface AIProviderBedrockSettings { * AccessKey. Write-only. */ readonly access_key_secret?: string; + /** + * RoleARN, when set, is the IAM role assumed via STS before calling + * Bedrock. The base identity (static keys or the AWS environment, e.g. + * IRSA / EKS Pod Identity / EC2 Instance Profile) signs the AssumeRole + * call, and the resulting temporary credentials sign Bedrock requests. + */ + readonly role_arn?: string; } // From codersdk/aiproviders_bedrock.go